fix: harden controlled AI session consistency

This commit is contained in:
2026-07-21 19:54:25 +08:00
parent 122f5d8c52
commit 251b212e06
15 changed files with 406 additions and 62 deletions

View File

@@ -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

View File

@@ -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")
}

View File

@@ -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
}

View 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)
}

View 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
}

View File

@@ -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

View 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"
}

View File

@@ -38,6 +38,7 @@ func AutoMigrate(database *gorm.DB) error {
&SenlinAgentProjectEvent{},
&SenlinAgentAIKey{},
&SenlinAgentAICallLog{},
&SenlinAgentAIRateBucket{},
&SenlinAgentTaskShare{},
)
}