fix: harden controlled AI session consistency
This commit is contained in:
@@ -18,6 +18,7 @@ import (
|
||||
type Gateway struct {
|
||||
systemKey string
|
||||
encryptionSecret string
|
||||
now func() time.Time
|
||||
}
|
||||
|
||||
type SelectedKey struct {
|
||||
@@ -36,7 +37,7 @@ func NewGateway(systemKey string) *Gateway {
|
||||
}
|
||||
|
||||
func NewGatewayWithSecret(systemKey string, encryptionSecret string) *Gateway {
|
||||
return &Gateway{systemKey: systemKey, encryptionSecret: encryptionSecret}
|
||||
return &Gateway{systemKey: systemKey, encryptionSecret: encryptionSecret, now: time.Now}
|
||||
}
|
||||
|
||||
func (g *Gateway) SaveUserKey(userID uint, provider string, apiKey string) error {
|
||||
@@ -55,11 +56,13 @@ func (g *Gateway) SelectKey(userID uint) (SelectedKey, error) {
|
||||
var userKey models.SenlinAgentAIKey
|
||||
err := models.DBService.Where("user_id = ?", userID).First(&userKey).Error
|
||||
if err == nil {
|
||||
selected := SelectedKey{Provider: userKey.Provider, KeyType: "user"}
|
||||
apiKey, err := decryptAPIKey(userKey.EncryptedAPIKey, g.encryptionSecret)
|
||||
if err != nil {
|
||||
return SelectedKey{}, err
|
||||
return selected, err
|
||||
}
|
||||
return SelectedKey{Provider: userKey.Provider, APIKey: apiKey, KeyType: "user"}, nil
|
||||
selected.APIKey = apiKey
|
||||
return selected, nil
|
||||
}
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return SelectedKey{}, err
|
||||
@@ -70,8 +73,11 @@ func (g *Gateway) SelectKey(userID uint) (SelectedKey, error) {
|
||||
return SelectedKey{Provider: "openai", APIKey: g.systemKey, KeyType: "system"}, nil
|
||||
}
|
||||
|
||||
func (g *Gateway) RecordCall(userID uint, provider string, usedKeyType string, action string, status string, errText string) error {
|
||||
return models.DBService.Create(&models.SenlinAgentAICallLog{
|
||||
func (g *Gateway) RecordCall(database *gorm.DB, userID uint, provider string, usedKeyType string, action string, status string, errText string) error {
|
||||
if database == nil {
|
||||
database = models.DBService
|
||||
}
|
||||
return database.Create(&models.SenlinAgentAICallLog{
|
||||
UserID: userID,
|
||||
Provider: provider,
|
||||
UsedKeyType: usedKeyType,
|
||||
@@ -81,17 +87,34 @@ func (g *Gateway) RecordCall(userID uint, provider string, usedKeyType string, a
|
||||
}).Error
|
||||
}
|
||||
|
||||
func (g *Gateway) CheckRateLimit(userID uint, action string, limit int, window time.Duration) error {
|
||||
// ReserveRateLimit 以数据库单条 UPSERT 原子占用固定窗口配额。
|
||||
// 配额在 provider/key 选择前占用,后续缺 key 或 provider 失败同样计入该窗口的尝试次数。
|
||||
func (g *Gateway) ReserveRateLimit(userID uint, action string, limit int, window time.Duration) error {
|
||||
if limit <= 0 {
|
||||
return nil
|
||||
}
|
||||
var count int64
|
||||
if err := models.DBService.Model(&models.SenlinAgentAICallLog{}).
|
||||
Where("user_id = ? AND action = ? AND created_at >= ?", userID, action, time.Now().Add(-window)).
|
||||
Count(&count).Error; err != nil {
|
||||
return err
|
||||
currentTime := time.Now().UTC()
|
||||
if g.now != nil {
|
||||
currentTime = g.now().UTC()
|
||||
}
|
||||
if count >= int64(limit) {
|
||||
windowStart := currentTime.Truncate(window)
|
||||
bucket := models.SenlinAgentAIRateBucket{
|
||||
UserID: userID, Action: action, WindowStart: windowStart, Count: 1,
|
||||
}
|
||||
result := models.DBService.Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "user_id"}, {Name: "action"}, {Name: "window_start"}},
|
||||
DoUpdates: clause.Assignments(map[string]any{
|
||||
"count": gorm.Expr("senlin_agent_ai_rate_buckets.count + 1"),
|
||||
"updated_at": currentTime,
|
||||
}),
|
||||
Where: clause.Where{Exprs: []clause.Expression{
|
||||
clause.Lt{Column: clause.Column{Table: "senlin_agent_ai_rate_buckets", Name: "count"}, Value: limit},
|
||||
}},
|
||||
}).Create(&bucket)
|
||||
if result.Error != nil {
|
||||
return result.Error
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return ErrAIRateLimited
|
||||
}
|
||||
return nil
|
||||
|
||||
@@ -63,7 +63,7 @@ func TestRecordCallStoresAuditFields(t *testing.T) {
|
||||
database := newTestDB(t)
|
||||
gateway := NewGateway("system-key")
|
||||
|
||||
require.NoError(t, gateway.RecordCall(3, "openai", "system", "inbox_analyze", "failed", "rate limited"))
|
||||
require.NoError(t, gateway.RecordCall(database, 3, "openai", "system", "inbox_analyze", "failed", "rate limited"))
|
||||
|
||||
var log models.SenlinAgentAICallLog
|
||||
require.NoError(t, database.First(&log).Error)
|
||||
@@ -75,12 +75,13 @@ func TestRecordCallStoresAuditFields(t *testing.T) {
|
||||
require.Equal(t, "rate limited", log.Error)
|
||||
}
|
||||
|
||||
func TestCheckRateLimitRejectsCallsOverWindow(t *testing.T) {
|
||||
newTestDB(t)
|
||||
func TestReserveRateLimitRejectsCallsOverWindow(t *testing.T) {
|
||||
database := newTestDB(t)
|
||||
require.NoError(t, database.Create(&models.SenlinAgentUser{Email: "rate@example.com", DisplayName: "Rate", PasswordHash: "hash"}).Error)
|
||||
gateway := NewGateway("system-key")
|
||||
require.NoError(t, gateway.RecordCall(3, "openai", "system", "inbox_analyze", "succeeded", ""))
|
||||
require.NoError(t, gateway.ReserveRateLimit(1, "inbox_analyze", 1, time.Hour))
|
||||
|
||||
err := gateway.CheckRateLimit(3, "inbox_analyze", 1, time.Hour)
|
||||
err := gateway.ReserveRateLimit(1, "inbox_analyze", 1, time.Hour)
|
||||
|
||||
require.ErrorContains(t, err, "ai rate limit exceeded")
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ package ai
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -73,7 +74,7 @@ func TestCreateAISessionReturnsRateLimitBeforeMissingKeyAndAuditsFailure(t *test
|
||||
project := createAIHandlerProject(t, database, owner.ID, "LIMITED")
|
||||
gateway := NewGatewayWithSecret("", "test-encryption-secret")
|
||||
for range aiSessionCreateLimit {
|
||||
require.NoError(t, gateway.RecordCall(owner.ID, "openai", "system", "ai_session_create", "ready", ""))
|
||||
require.NoError(t, gateway.ReserveRateLimit(owner.ID, aiSessionCreateAction, aiSessionCreateLimit, time.Hour))
|
||||
}
|
||||
router := aiHandlerTestRouter(owner.ID, gateway)
|
||||
recorder := httptest.NewRecorder()
|
||||
@@ -130,6 +131,69 @@ func TestCreateAISessionWithoutKeyReturnsAuditedErrorAndCreatesNoFormalObjects(t
|
||||
require.Equal(t, "ai_session_create", call.Action)
|
||||
require.Equal(t, "failed", call.Status)
|
||||
require.Equal(t, "ai_key_missing", call.Error)
|
||||
var bucket models.SenlinAgentAIRateBucket
|
||||
require.NoError(t, database.Where("user_id = ? AND action = ?", owner.ID, aiSessionCreateAction).First(&bucket).Error)
|
||||
require.Equal(t, 1, bucket.Count)
|
||||
}
|
||||
|
||||
func TestCreateAISessionRollsBackSessionWhenReadyAuditWriteFails(t *testing.T) {
|
||||
database := newAIHandlerTestDB(t)
|
||||
owner := createAIHandlerUser(t, database, "audit-failure@example.com")
|
||||
project := createAIHandlerProject(t, database, owner.ID, "AUDIT_FAILURE")
|
||||
injectedError := errors.New("injected ready audit failure")
|
||||
callbackName := "test:fail_ready_ai_audit"
|
||||
require.NoError(t, database.Callback().Create().Before("gorm:create").Register(callbackName, func(tx *gorm.DB) {
|
||||
call, ok := tx.Statement.Dest.(*models.SenlinAgentAICallLog)
|
||||
if ok && call.Status == defaultSessionStatus {
|
||||
tx.AddError(injectedError)
|
||||
}
|
||||
}))
|
||||
t.Cleanup(func() { database.Callback().Create().Remove(callbackName) })
|
||||
router := aiHandlerTestRouter(owner.ID, NewGatewayWithSecret("system-key", "test-encryption-secret"))
|
||||
recorder := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(recorder, authenticatedAIRequest(t, http.MethodPost, "/api/v1/projects/"+project.Identity+"/ai-sessions", map[string]any{
|
||||
"title": "必须回滚", "context": "成功审计失败时不能残留会话",
|
||||
}))
|
||||
|
||||
require.Equal(t, http.StatusInternalServerError, recorder.Code)
|
||||
var sessionCount int64
|
||||
require.NoError(t, database.Model(&models.SenlinAgentAISession{}).Where("project_id = ?", project.ID).Count(&sessionCount).Error)
|
||||
require.Zero(t, sessionCount)
|
||||
var calls []models.SenlinAgentAICallLog
|
||||
require.NoError(t, database.Where("user_id = ? AND action = ?", owner.ID, aiSessionCreateAction).Find(&calls).Error)
|
||||
require.Len(t, calls, 1)
|
||||
require.Equal(t, "openai", calls[0].Provider)
|
||||
require.Equal(t, "system", calls[0].UsedKeyType)
|
||||
require.Equal(t, "failed", calls[0].Status)
|
||||
require.Equal(t, "audit_write_failed", calls[0].Error)
|
||||
}
|
||||
|
||||
func TestCreateAISessionAuditsKnownProviderMetadataWhenUserKeyDecryptFails(t *testing.T) {
|
||||
database := newAIHandlerTestDB(t)
|
||||
owner := createAIHandlerUser(t, database, "decrypt-failure@example.com")
|
||||
project := createAIHandlerProject(t, database, owner.ID, "DECRYPT_FAILURE")
|
||||
require.NoError(t, database.Create(&models.SenlinAgentAIKey{
|
||||
UserID: owner.ID, Provider: "deepseek", EncryptedAPIKey: "v1:not-valid-base64",
|
||||
}).Error)
|
||||
router := aiHandlerTestRouter(owner.ID, NewGatewayWithSecret("system-key", "test-encryption-secret"))
|
||||
recorder := httptest.NewRecorder()
|
||||
|
||||
router.ServeHTTP(recorder, authenticatedAIRequest(t, http.MethodPost, "/api/v1/projects/"+project.Identity+"/ai-sessions", map[string]any{
|
||||
"title": "解密失败", "context": "审计不得丢失已知元数据",
|
||||
}))
|
||||
|
||||
require.Equal(t, http.StatusInternalServerError, recorder.Code)
|
||||
var call models.SenlinAgentAICallLog
|
||||
require.NoError(t, database.Where("user_id = ? AND action = ?", owner.ID, aiSessionCreateAction).First(&call).Error)
|
||||
require.Equal(t, "deepseek", call.Provider)
|
||||
require.Equal(t, "user", call.UsedKeyType)
|
||||
require.Equal(t, "failed", call.Status)
|
||||
require.Equal(t, "provider_selection_failed", call.Error)
|
||||
require.NotContains(t, call.Error, "not-valid-base64")
|
||||
var sessionCount int64
|
||||
require.NoError(t, database.Model(&models.SenlinAgentAISession{}).Count(&sessionCount).Error)
|
||||
require.Zero(t, sessionCount)
|
||||
}
|
||||
|
||||
func TestCreateAISessionReturnsIdentityDTOAndCompleteAuditWithoutAutomaticObjectIDs(t *testing.T) {
|
||||
@@ -205,7 +269,7 @@ type recordingSessionGateway struct {
|
||||
selectErr error
|
||||
}
|
||||
|
||||
func (g *recordingSessionGateway) CheckRateLimit(uint, string, int, time.Duration) error {
|
||||
func (g *recordingSessionGateway) ReserveRateLimit(uint, string, int, time.Duration) error {
|
||||
g.steps = append(g.steps, "rate")
|
||||
return g.rateErr
|
||||
}
|
||||
@@ -215,7 +279,7 @@ func (g *recordingSessionGateway) SelectKey(uint) (SelectedKey, error) {
|
||||
return g.selected, g.selectErr
|
||||
}
|
||||
|
||||
func (g *recordingSessionGateway) RecordCall(uint, string, string, string, string, string) error {
|
||||
func (g *recordingSessionGateway) RecordCall(*gorm.DB, uint, string, string, string, string, string) error {
|
||||
g.steps = append(g.steps, "record")
|
||||
return nil
|
||||
}
|
||||
|
||||
72
backend/internal/logic/ai/rate_limit_postgres_test.go
Normal file
72
backend/internal/logic/ai/rate_limit_postgres_test.go
Normal file
@@ -0,0 +1,72 @@
|
||||
package ai
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/driver/postgres"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
"senlinai-agent/backend/internal/models"
|
||||
)
|
||||
|
||||
func TestPostgresReserveRateLimitIsAtomicAcrossConcurrentConnections(t *testing.T) {
|
||||
dsn := os.Getenv("DATABASE_URL")
|
||||
if dsn == "" {
|
||||
t.Skip("DATABASE_URL is not configured; skipping PostgreSQL AI rate reservation test")
|
||||
}
|
||||
database, err := gorm.Open(postgres.Open(dsn), &gorm.Config{TranslateError: true, Logger: logger.Default.LogMode(logger.Silent)})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, models.AutoMigrate(database))
|
||||
models.DBService = database
|
||||
suffix := fmt.Sprint(time.Now().UnixNano())
|
||||
user := createAIRateTestUser(t, database, "postgres-rate-"+suffix+"@example.com")
|
||||
action := "postgres_concurrent_" + suffix
|
||||
t.Cleanup(func() {
|
||||
database.Where("user_id = ?", user.ID).Delete(&models.SenlinAgentAIRateBucket{})
|
||||
database.Delete(&user)
|
||||
})
|
||||
gateway := NewGatewayWithSecret("system-key", "test-encryption-secret")
|
||||
const (
|
||||
limit = 9
|
||||
attempts = 48
|
||||
)
|
||||
|
||||
start := make(chan struct{})
|
||||
results := make(chan error, attempts)
|
||||
var wait sync.WaitGroup
|
||||
for range attempts {
|
||||
wait.Add(1)
|
||||
go func() {
|
||||
defer wait.Done()
|
||||
<-start
|
||||
results <- gateway.ReserveRateLimit(user.ID, action, limit, time.Hour)
|
||||
}()
|
||||
}
|
||||
close(start)
|
||||
wait.Wait()
|
||||
close(results)
|
||||
|
||||
allowed := 0
|
||||
limited := 0
|
||||
for err := range results {
|
||||
switch {
|
||||
case err == nil:
|
||||
allowed++
|
||||
case errors.Is(err, ErrAIRateLimited):
|
||||
limited++
|
||||
default:
|
||||
require.NoError(t, err)
|
||||
}
|
||||
}
|
||||
require.Equal(t, limit, allowed)
|
||||
require.Equal(t, attempts-limit, limited)
|
||||
var bucket models.SenlinAgentAIRateBucket
|
||||
require.NoError(t, database.Where("user_id = ? AND action = ?", user.ID, action).First(&bucket).Error)
|
||||
require.Equal(t, limit, bucket.Count)
|
||||
}
|
||||
101
backend/internal/logic/ai/rate_limit_test.go
Normal file
101
backend/internal/logic/ai/rate_limit_test.go
Normal file
@@ -0,0 +1,101 @@
|
||||
package ai
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
"senlinai-agent/backend/internal/models"
|
||||
)
|
||||
|
||||
func TestReserveRateLimitIsAtomicUnderConcurrentSQLiteRequests(t *testing.T) {
|
||||
database := newConcurrentAIRateTestDB(t)
|
||||
user := createAIRateTestUser(t, database, "sqlite-rate@example.com")
|
||||
gateway := NewGatewayWithSecret("system-key", "test-encryption-secret")
|
||||
const (
|
||||
limit = 7
|
||||
attempts = 40
|
||||
)
|
||||
|
||||
start := make(chan struct{})
|
||||
results := make(chan error, attempts)
|
||||
var wait sync.WaitGroup
|
||||
for range attempts {
|
||||
wait.Add(1)
|
||||
go func() {
|
||||
defer wait.Done()
|
||||
<-start
|
||||
results <- gateway.ReserveRateLimit(user.ID, "concurrent_session_create", limit, time.Hour)
|
||||
}()
|
||||
}
|
||||
close(start)
|
||||
wait.Wait()
|
||||
close(results)
|
||||
|
||||
allowed := 0
|
||||
limited := 0
|
||||
for err := range results {
|
||||
switch {
|
||||
case err == nil:
|
||||
allowed++
|
||||
case errors.Is(err, ErrAIRateLimited):
|
||||
limited++
|
||||
default:
|
||||
require.NoError(t, err)
|
||||
}
|
||||
}
|
||||
require.Equal(t, limit, allowed)
|
||||
require.Equal(t, attempts-limit, limited)
|
||||
var buckets []models.SenlinAgentAIRateBucket
|
||||
require.NoError(t, database.Where("user_id = ? AND action = ?", user.ID, "concurrent_session_create").Find(&buckets).Error)
|
||||
require.Len(t, buckets, 1)
|
||||
require.Equal(t, limit, buckets[0].Count)
|
||||
}
|
||||
|
||||
func TestReserveRateLimitUsesFixedWindowsAndCountsFailedAttempts(t *testing.T) {
|
||||
database := newConcurrentAIRateTestDB(t)
|
||||
user := createAIRateTestUser(t, database, "window-rate@example.com")
|
||||
gateway := NewGatewayWithSecret("system-key", "test-encryption-secret")
|
||||
current := time.Date(2026, 7, 21, 10, 15, 0, 0, time.UTC)
|
||||
gateway.now = func() time.Time { return current }
|
||||
|
||||
require.NoError(t, gateway.ReserveRateLimit(user.ID, "windowed_session_create", 2, time.Hour))
|
||||
require.NoError(t, gateway.ReserveRateLimit(user.ID, "windowed_session_create", 2, time.Hour))
|
||||
require.ErrorIs(t, gateway.ReserveRateLimit(user.ID, "windowed_session_create", 2, time.Hour), ErrAIRateLimited)
|
||||
|
||||
current = current.Add(time.Hour)
|
||||
require.NoError(t, gateway.ReserveRateLimit(user.ID, "windowed_session_create", 2, time.Hour))
|
||||
var buckets []models.SenlinAgentAIRateBucket
|
||||
require.NoError(t, database.Where("user_id = ? AND action = ?", user.ID, "windowed_session_create").Order("window_start asc").Find(&buckets).Error)
|
||||
require.Len(t, buckets, 2)
|
||||
require.Equal(t, []int{2, 1}, []int{buckets[0].Count, buckets[1].Count})
|
||||
}
|
||||
|
||||
func newConcurrentAIRateTestDB(t *testing.T) *gorm.DB {
|
||||
t.Helper()
|
||||
databasePath := filepath.ToSlash(filepath.Join(t.TempDir(), "ai-rate.db"))
|
||||
dsn := fmt.Sprintf("file:%s?_pragma=busy_timeout(10000)&_pragma=journal_mode(WAL)", databasePath)
|
||||
database, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
|
||||
require.NoError(t, err)
|
||||
sqlDatabase, err := database.DB()
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { require.NoError(t, sqlDatabase.Close()) })
|
||||
sqlDatabase.SetMaxOpenConns(20)
|
||||
require.NoError(t, models.AutoMigrate(database))
|
||||
models.DBService = database
|
||||
return database
|
||||
}
|
||||
|
||||
func createAIRateTestUser(t *testing.T, database *gorm.DB, email string) models.SenlinAgentUser {
|
||||
t.Helper()
|
||||
user := models.SenlinAgentUser{Email: email, DisplayName: email, PasswordHash: "hash"}
|
||||
require.NoError(t, database.Create(&user).Error)
|
||||
return user
|
||||
}
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"senlinai-agent/backend/internal/logic/projects"
|
||||
"senlinai-agent/backend/internal/models"
|
||||
)
|
||||
@@ -21,9 +22,9 @@ var (
|
||||
)
|
||||
|
||||
type sessionGateway interface {
|
||||
CheckRateLimit(userID uint, action string, limit int, window time.Duration) error
|
||||
ReserveRateLimit(userID uint, action string, limit int, window time.Duration) error
|
||||
SelectKey(userID uint) (SelectedKey, error)
|
||||
RecordCall(userID uint, provider string, usedKeyType string, action string, status string, errText string) error
|
||||
RecordCall(database *gorm.DB, userID uint, provider string, usedKeyType string, action string, status string, errText string) error
|
||||
}
|
||||
|
||||
// SessionService 只管理项目内普通 AI 会话;它不会把上下文自动转换为任务、笔记或资料。
|
||||
@@ -66,14 +67,14 @@ func (s *SessionService) Create(userID uint, projectIdentity, title, context str
|
||||
return nil, errors.New("ai gateway is required")
|
||||
}
|
||||
|
||||
if err := s.gateway.CheckRateLimit(userID, aiSessionCreateAction, aiSessionCreateLimit, time.Hour); err != nil {
|
||||
if err := s.gateway.ReserveRateLimit(userID, aiSessionCreateAction, aiSessionCreateLimit, time.Hour); err != nil {
|
||||
if errors.Is(err, ErrAIRateLimited) {
|
||||
if auditErr := s.gateway.RecordCall(userID, "none", "none", aiSessionCreateAction, "failed", "ai_rate_limited"); auditErr != nil {
|
||||
if auditErr := s.gateway.RecordCall(models.DBService, userID, "none", "none", aiSessionCreateAction, "failed", "ai_rate_limited"); auditErr != nil {
|
||||
return nil, fmt.Errorf("record ai rate limit failure: %w", auditErr)
|
||||
}
|
||||
return nil, ErrAIRateLimited
|
||||
}
|
||||
if auditErr := s.gateway.RecordCall(userID, "none", "none", aiSessionCreateAction, "failed", "rate_limit_check_failed"); auditErr != nil {
|
||||
if auditErr := s.gateway.RecordCall(models.DBService, userID, "none", "none", aiSessionCreateAction, "failed", "rate_limit_reservation_failed"); auditErr != nil {
|
||||
return nil, fmt.Errorf("check rate limit: %v; record failure: %w", err, auditErr)
|
||||
}
|
||||
return nil, err
|
||||
@@ -85,7 +86,8 @@ func (s *SessionService) Create(userID uint, projectIdentity, title, context str
|
||||
if errors.Is(err, ErrAIKeyMissing) {
|
||||
code = "ai_key_missing"
|
||||
}
|
||||
if auditErr := s.gateway.RecordCall(userID, "none", "none", aiSessionCreateAction, "failed", code); auditErr != nil {
|
||||
provider, keyType := selectedAuditMetadata(selected)
|
||||
if auditErr := s.gateway.RecordCall(models.DBService, userID, provider, keyType, aiSessionCreateAction, "failed", code); auditErr != nil {
|
||||
return nil, fmt.Errorf("select ai key: %v; record failure: %w", err, auditErr)
|
||||
}
|
||||
return nil, err
|
||||
@@ -98,19 +100,40 @@ func (s *SessionService) Create(userID uint, projectIdentity, title, context str
|
||||
Context: context,
|
||||
Status: defaultSessionStatus,
|
||||
}
|
||||
if err := models.DBService.Create(&session).Error; err != nil {
|
||||
if auditErr := s.gateway.RecordCall(userID, selected.Provider, selected.KeyType, aiSessionCreateAction, "failed", "session_create_failed"); auditErr != nil {
|
||||
return nil, fmt.Errorf("create ai session: %v; record failure: %w", err, auditErr)
|
||||
failureCode := "session_create_failed"
|
||||
err = models.DBService.Transaction(func(tx *gorm.DB) error {
|
||||
if err := tx.Create(&session).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
// ready 只表示会话入口已建立,不表示 provider 已回复或任何业务对象已创建。
|
||||
failureCode = "audit_write_failed"
|
||||
if err := s.gateway.RecordCall(tx, userID, selected.Provider, selected.KeyType, aiSessionCreateAction, defaultSessionStatus, ""); err != nil {
|
||||
return err
|
||||
}
|
||||
failureCode = "session_transaction_failed"
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
if auditErr := s.gateway.RecordCall(models.DBService, userID, selected.Provider, selected.KeyType, aiSessionCreateAction, "failed", failureCode); auditErr != nil {
|
||||
return nil, fmt.Errorf("create ai session transaction: %v; record failure: %w", err, auditErr)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
// ready 只表示会话入口已建立,不表示 provider 已回复或任何业务对象已创建。
|
||||
if err := s.gateway.RecordCall(userID, selected.Provider, selected.KeyType, aiSessionCreateAction, defaultSessionStatus, ""); err != nil {
|
||||
return nil, fmt.Errorf("record ai session creation: %w", err)
|
||||
}
|
||||
return &session, nil
|
||||
}
|
||||
|
||||
func selectedAuditMetadata(selected SelectedKey) (string, string) {
|
||||
provider := strings.TrimSpace(selected.Provider)
|
||||
keyType := strings.TrimSpace(selected.KeyType)
|
||||
if provider == "" {
|
||||
provider = "none"
|
||||
}
|
||||
if keyType == "" {
|
||||
keyType = "none"
|
||||
}
|
||||
return provider, keyType
|
||||
}
|
||||
|
||||
func aiSessionStatus(session models.SenlinAgentAISession) string {
|
||||
if status := strings.TrimSpace(session.Status); status != "" {
|
||||
return status
|
||||
|
||||
19
backend/internal/models/ai_rate_bucket.go
Normal file
19
backend/internal/models/ai_rate_bucket.go
Normal file
@@ -0,0 +1,19 @@
|
||||
package models
|
||||
|
||||
import "time"
|
||||
|
||||
// SenlinAgentAIRateBucket 保存用户在固定窗口内已占用的 AI 请求配额。
|
||||
// 复合唯一键让单条 UPSERT 在数据库层完成跨进程并发仲裁。
|
||||
type SenlinAgentAIRateBucket struct {
|
||||
ID uint `gorm:"primaryKey"`
|
||||
UserID uint `gorm:"not null;uniqueIndex:uidx_senlin_agent_ai_rate_bucket,priority:1"`
|
||||
Action string `gorm:"size:100;not null;uniqueIndex:uidx_senlin_agent_ai_rate_bucket,priority:2"`
|
||||
WindowStart time.Time `gorm:"not null;uniqueIndex:uidx_senlin_agent_ai_rate_bucket,priority:3"`
|
||||
Count int `gorm:"not null"`
|
||||
CreatedAt time.Time
|
||||
UpdatedAt time.Time
|
||||
}
|
||||
|
||||
func (SenlinAgentAIRateBucket) TableName() string {
|
||||
return "senlin_agent_ai_rate_buckets"
|
||||
}
|
||||
@@ -38,6 +38,7 @@ func AutoMigrate(database *gorm.DB) error {
|
||||
&SenlinAgentProjectEvent{},
|
||||
&SenlinAgentAIKey{},
|
||||
&SenlinAgentAICallLog{},
|
||||
&SenlinAgentAIRateBucket{},
|
||||
&SenlinAgentTaskShare{},
|
||||
)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user