Files
PR-Helper/handlers/auth.go
T

161 lines
4.4 KiB
Go

package handlers
import (
"database/sql"
"net/http"
"strings"
"time"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"golang.org/x/crypto/bcrypt"
)
type AuthHandler struct {
db *sql.DB
}
func NewAuthHandler(db *sql.DB) *AuthHandler {
return &AuthHandler{db: db}
}
// Login renders the login page.
func (h *AuthHandler) Login(c *gin.Context) {
c.HTML(http.StatusOK, "pages/login.html", gin.H{})
}
// Register renders the registration page.
func (h *AuthHandler) Register(c *gin.Context) {
c.HTML(http.StatusOK, "pages/register.html", gin.H{})
}
// HandleLogin processes POST /api/auth/login.
func (h *AuthHandler) HandleLogin(c *gin.Context) {
var req struct {
Email string `json:"email" binding:"required"`
Password string `json:"password" binding:"required"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "邮箱和密码为必填项"})
return
}
req.Email = strings.TrimSpace(strings.ToLower(req.Email))
var user struct {
ID int64
PasswordHash string
}
err := h.db.QueryRow(`SELECT id, password_hash FROM users WHERE email = ?`, req.Email).
Scan(&user.ID, &user.PasswordHash)
if err == sql.ErrNoRows {
c.JSON(http.StatusUnauthorized, gin.H{"error": "邮箱或密码错误"})
return
}
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "服务器错误"})
return
}
if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(req.Password)); err != nil {
c.JSON(http.StatusUnauthorized, gin.H{"error": "邮箱或密码错误"})
return
}
// Set session
session := sessions.Default(c)
session.Set("user_id", user.ID)
if err := session.Save(); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "会话保存失败"})
return
}
c.JSON(http.StatusOK, gin.H{"ok": true})
}
// HandleRegister processes POST /api/auth/register.
func (h *AuthHandler) HandleRegister(c *gin.Context) {
var req struct {
Email string `json:"email" binding:"required"`
Password string `json:"password" binding:"required"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": "邮箱和密码为必填项"})
return
}
req.Email = strings.TrimSpace(strings.ToLower(req.Email))
// Validate email format (basic check)
if !strings.Contains(req.Email, "@") || !strings.Contains(req.Email, ".") {
c.JSON(http.StatusBadRequest, gin.H{"error": "邮箱格式不正确"})
return
}
// Validate password length
if len(req.Password) < 6 {
c.JSON(http.StatusBadRequest, gin.H{"error": "密码长度至少为 6 位"})
return
}
// Check if email already exists
var exists int
h.db.QueryRow(`SELECT COUNT(*) FROM users WHERE email = ?`, req.Email).Scan(&exists)
if exists > 0 {
c.JSON(http.StatusConflict, gin.H{"error": "该邮箱已被注册"})
return
}
// Hash password
hash, err := bcrypt.GenerateFromPassword([]byte(req.Password), bcrypt.DefaultCost)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "密码加密失败"})
return
}
// Insert user
now := time.Now().UTC()
result, err := h.db.Exec(`INSERT INTO users (email, password_hash, created_at, updated_at) VALUES (?, ?, ?, ?)`,
req.Email, string(hash), now, now)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "注册失败"})
return
}
userID, _ := result.LastInsertId()
// Seed default settings for the user
for key, val := range defaultUserSettings {
h.db.Exec("INSERT IGNORE INTO user_settings (user_id, `key`, value) VALUES (?, ?, ?)", userID, key, val)
}
// Auto-login: set session
session := sessions.Default(c)
session.Set("user_id", userID)
if err := session.Save(); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "会话保存失败"})
return
}
c.JSON(http.StatusOK, gin.H{"ok": true})
}
// HandleLogout processes POST /api/auth/logout.
func (h *AuthHandler) HandleLogout(c *gin.Context) {
session := sessions.Default(c)
session.Clear()
session.Save()
c.Redirect(http.StatusFound, "/login")
}
// defaultUserSettings mirrors models.DefaultSettings for seeding new users.
var defaultUserSettings = map[string]string{
"llm.endpoint": "https://api.deepseek.com",
"llm.api_key": "",
"llm.model": "deepseek-v4-pro",
"review.top_n": "20",
"review.concurrency": "5",
"cache.max_age_days": "7",
"cache.max_size_mb": "5000",
}