refactor backend packages and global db service
This commit is contained in:
37
backend/internal/logic/auth/middleware.go
Normal file
37
backend/internal/logic/auth/middleware.go
Normal file
@@ -0,0 +1,37 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
const CurrentUserIDKey = "currentUserID"
|
||||
|
||||
func RequireUser(tokenVerifier func(string) (uint, error)) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
header := c.GetHeader("Authorization")
|
||||
token := strings.TrimPrefix(header, "Bearer ")
|
||||
if token == "" || token == header {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "missing bearer token"})
|
||||
return
|
||||
}
|
||||
userID, err := tokenVerifier(token)
|
||||
if err != nil {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "invalid bearer token"})
|
||||
return
|
||||
}
|
||||
c.Set(CurrentUserIDKey, userID)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
func CurrentUserID(c *gin.Context) (uint, bool) {
|
||||
value, ok := c.Get(CurrentUserIDKey)
|
||||
if !ok {
|
||||
return 0, false
|
||||
}
|
||||
userID, ok := value.(uint)
|
||||
return userID, ok && userID > 0
|
||||
}
|
||||
136
backend/internal/logic/auth/service.go
Normal file
136
backend/internal/logic/auth/service.go
Normal file
@@ -0,0 +1,136 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
"senlinai-agent/backend/internal/models"
|
||||
)
|
||||
|
||||
type Service struct {
|
||||
secret string
|
||||
inviteTokenTTL time.Duration
|
||||
sessionTTL time.Duration
|
||||
}
|
||||
|
||||
func NewService(secret string) *Service {
|
||||
return &Service{secret: secret, inviteTokenTTL: 7 * 24 * time.Hour, sessionTTL: 24 * time.Hour}
|
||||
}
|
||||
|
||||
func (s *Service) CreateInvite(adminID uint, email string) (string, error) {
|
||||
normalized := strings.ToLower(strings.TrimSpace(email))
|
||||
if normalized == "" {
|
||||
return "", errors.New("email is required")
|
||||
}
|
||||
return s.signToken(tokenPayload{
|
||||
Type: "invite",
|
||||
Subject: normalized,
|
||||
IssuerID: adminID,
|
||||
ExpiresAt: time.Now().Add(s.inviteTokenTTL).Unix(),
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Service) RegisterWithInvite(token, email, displayName, password string) (*models.SenlinAgentUser, error) {
|
||||
payload, err := s.verifyToken(token, "invite")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if payload.Subject != strings.ToLower(strings.TrimSpace(email)) {
|
||||
return nil, errors.New("invite email mismatch")
|
||||
}
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
user := &models.SenlinAgentUser{
|
||||
Email: payload.Subject,
|
||||
DisplayName: strings.TrimSpace(displayName),
|
||||
PasswordHash: string(hash),
|
||||
Role: "user",
|
||||
}
|
||||
return user, models.DBService.Create(user).Error
|
||||
}
|
||||
|
||||
func (s *Service) Login(email, password string) (string, error) {
|
||||
var user models.SenlinAgentUser
|
||||
if err := models.DBService.Where("email = ?", strings.ToLower(strings.TrimSpace(email))).First(&user).Error; err != nil {
|
||||
return "", errors.New("invalid credentials")
|
||||
}
|
||||
if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(password)); err != nil {
|
||||
return "", errors.New("invalid credentials")
|
||||
}
|
||||
return s.signToken(tokenPayload{
|
||||
Type: "session",
|
||||
Subject: fmt.Sprintf("%d", user.ID),
|
||||
ExpiresAt: time.Now().Add(s.sessionTTL).Unix(),
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Service) VerifySession(token string) (uint, error) {
|
||||
payload, err := s.verifyToken(token, "session")
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
var userID uint
|
||||
if _, err := fmt.Sscanf(payload.Subject, "%d", &userID); err != nil || userID == 0 {
|
||||
return 0, errors.New("invalid session subject")
|
||||
}
|
||||
return userID, nil
|
||||
}
|
||||
|
||||
type tokenPayload struct {
|
||||
Type string `json:"type"`
|
||||
Subject string `json:"subject"`
|
||||
IssuerID uint `json:"issuer_id,omitempty"`
|
||||
ExpiresAt int64 `json:"expires_at"`
|
||||
}
|
||||
|
||||
func (s *Service) signToken(payload tokenPayload) (string, error) {
|
||||
body, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
encoded := base64.RawURLEncoding.EncodeToString(body)
|
||||
return encoded + "." + s.sign(encoded), nil
|
||||
}
|
||||
|
||||
func (s *Service) verifyToken(token string, expectedType string) (tokenPayload, error) {
|
||||
separator := strings.LastIndex(token, ".")
|
||||
if separator < 1 || separator == len(token)-1 {
|
||||
return tokenPayload{}, errors.New("invalid token")
|
||||
}
|
||||
encoded := token[:separator]
|
||||
signature := token[separator+1:]
|
||||
if !hmac.Equal([]byte(signature), []byte(s.sign(encoded))) {
|
||||
return tokenPayload{}, errors.New("invalid token signature")
|
||||
}
|
||||
body, err := base64.RawURLEncoding.DecodeString(encoded)
|
||||
if err != nil {
|
||||
return tokenPayload{}, errors.New("invalid token payload")
|
||||
}
|
||||
var payload tokenPayload
|
||||
if err := json.Unmarshal(body, &payload); err != nil {
|
||||
return tokenPayload{}, errors.New("invalid token payload")
|
||||
}
|
||||
if payload.Type != expectedType {
|
||||
return tokenPayload{}, errors.New("invalid token type")
|
||||
}
|
||||
if payload.ExpiresAt <= time.Now().Unix() {
|
||||
return tokenPayload{}, errors.New("token expired")
|
||||
}
|
||||
return payload, nil
|
||||
}
|
||||
|
||||
func (s *Service) sign(value string) string {
|
||||
mac := hmac.New(sha256.New, []byte(s.secret))
|
||||
mac.Write([]byte(value))
|
||||
return hex.EncodeToString(mac.Sum(nil))
|
||||
}
|
||||
84
backend/internal/logic/auth/service_test.go
Normal file
84
backend/internal/logic/auth/service_test.go
Normal file
@@ -0,0 +1,84 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
"senlinai-agent/backend/internal/models"
|
||||
)
|
||||
|
||||
func TestRegisterWithInviteCreatesUser(t *testing.T) {
|
||||
newTestDB(t)
|
||||
service := NewService("test-secret")
|
||||
token, err := service.CreateInvite(1, "lead@example.com")
|
||||
require.NoError(t, err)
|
||||
|
||||
user, err := service.RegisterWithInvite(token, "lead@example.com", "Lead", "password123")
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "lead@example.com", user.Email)
|
||||
require.NotEmpty(t, user.PasswordHash)
|
||||
require.NotEqual(t, "password123", user.PasswordHash)
|
||||
}
|
||||
|
||||
func TestRegisterWithInviteRejectsWrongEmail(t *testing.T) {
|
||||
newTestDB(t)
|
||||
service := NewService("test-secret")
|
||||
token, err := service.CreateInvite(1, "lead@example.com")
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = service.RegisterWithInvite(token, "other@example.com", "Other", "password123")
|
||||
|
||||
require.ErrorContains(t, err, "invite email mismatch")
|
||||
}
|
||||
|
||||
func TestLoginRejectsInvalidPassword(t *testing.T) {
|
||||
newTestDB(t)
|
||||
service := NewService("test-secret")
|
||||
token, err := service.CreateInvite(1, "lead@example.com")
|
||||
require.NoError(t, err)
|
||||
_, err = service.RegisterWithInvite(token, "lead@example.com", "Lead", "password123")
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = service.Login("lead@example.com", "wrong-password")
|
||||
|
||||
require.ErrorContains(t, err, "invalid credentials")
|
||||
}
|
||||
|
||||
func TestLoginReturnsSignedSessionToken(t *testing.T) {
|
||||
newTestDB(t)
|
||||
service := NewService("test-secret")
|
||||
token, err := service.CreateInvite(1, "lead@example.com")
|
||||
require.NoError(t, err)
|
||||
user, err := service.RegisterWithInvite(token, "lead@example.com", "Lead", "password123")
|
||||
require.NoError(t, err)
|
||||
|
||||
sessionToken, err := service.Login("lead@example.com", "password123")
|
||||
require.NoError(t, err)
|
||||
userID, err := service.VerifySession(sessionToken)
|
||||
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, user.ID, userID)
|
||||
require.NotEqual(t, fmt.Sprintf("user:%d", user.ID), sessionToken)
|
||||
}
|
||||
|
||||
func TestVerifySessionRejectsForgeableLegacyToken(t *testing.T) {
|
||||
newTestDB(t)
|
||||
service := NewService("test-secret")
|
||||
|
||||
_, err := service.VerifySession("user:1")
|
||||
|
||||
require.ErrorContains(t, err, "invalid token")
|
||||
}
|
||||
|
||||
func newTestDB(t *testing.T) *gorm.DB {
|
||||
t.Helper()
|
||||
database, err := gorm.Open(sqlite.Open(fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())), &gorm.Config{})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, models.AutoMigrate(database))
|
||||
models.DBService = database
|
||||
return database
|
||||
}
|
||||
Reference in New Issue
Block a user