diff --git a/apps/web_v1/scripts/structure-check.mjs b/apps/web_v1/scripts/structure-check.mjs
index 2097376..23ebb0d 100644
--- a/apps/web_v1/scripts/structure-check.mjs
+++ b/apps/web_v1/scripts/structure-check.mjs
@@ -106,12 +106,17 @@ const aiApiSource = readFileSync('src/api/ai.ts', 'utf8')
for (const required of ['/api/v1/projects/', '/ai-sessions', 'listAISessions', 'createAISession']) {
if (!aiApiSource.includes(required)) failures.push(`AI API must include ${required}`)
}
+if (!aiApiSource.includes('signal?: AbortSignal')) failures.push('AI API requests must accept an AbortSignal')
+if (!aiApiSource.includes('signal,')) failures.push('AI API requests must pass the AbortSignal to apiRequest')
if (/\b(?:task|note|source)Id\b/.test(aiApiSource)) failures.push('AI session responses must not expose automatic formal object IDs')
const aiPageSource = readFileSync('src/pages/projects/project-ai.tsx', 'utf8')
for (const required of ['AI 助手', '创建会话', 'loading', 'error']) {
if (!aiPageSource.includes(required)) failures.push(`project AI page must include ${required}`)
}
+for (const required of ['AbortController', 'generationRef', 'projectRef', 'setSessions([])']) {
+ if (!aiPageSource.includes(required)) failures.push(`project AI page must gate stale requests with ${required}`)
+}
for (const forbidden of ['DeepSeek V4.0 Flash', '给 DeepSeek 发送消息', 'IconAttachment', 'agent-send-button']) {
if (aiPageSource.includes(forbidden)) failures.push(`project AI page contains unsupported chat control ${forbidden}`)
}
@@ -119,7 +124,7 @@ for (const forbidden of ['DeepSeek V4.0 Flash', '给 DeepSeek 发送消息', 'Ic
const unsupportedControls = [
{ file: 'src/pages/workspace-explore.tsx', required: '暂未开放', forbidden: ['同步数据源', '添加数据源', 'IconRefresh', 'IconEdit', 'IconDelete'] },
{ file: 'src/pages/projects/project-new-channel.tsx', required: '暂未开放', forbidden: ['保存频道', '(`/api/v1/projects/${projectId}/ai-sessions`, {
token: session.token,
+ signal,
})
}
-export async function createAISession(session: ApiSession, projectId: string, input: CreateAISessionInput) {
+export async function createAISession(session: ApiSession, projectId: string, input: CreateAISessionInput, signal?: AbortSignal) {
return apiRequest(`/api/v1/projects/${projectId}/ai-sessions`, {
method: 'POST',
token: session.token,
body: input,
+ signal,
})
}
diff --git a/apps/web_v1/src/app/App.tsx b/apps/web_v1/src/app/App.tsx
index 31bbe1a..ab065ad 100644
--- a/apps/web_v1/src/app/App.tsx
+++ b/apps/web_v1/src/app/App.tsx
@@ -44,14 +44,14 @@ function App() {
const [searchResultPreview, setSearchResultPreview] = useState(null)
const workspaceSearch = useWorkbenchSearch(session)
- const handleListAISessions = useCallback((projectId: string) => {
+ const handleListAISessions = useCallback((projectId: string, signal?: AbortSignal) => {
if (!session) return Promise.reject(new Error('未登录'))
- return listAISessions(session, projectId)
+ return listAISessions(session, projectId, signal)
}, [session])
- const handleCreateAISession = useCallback((projectId: string, input: CreateAISessionInput) => {
+ const handleCreateAISession = useCallback((projectId: string, input: CreateAISessionInput, signal?: AbortSignal) => {
if (!session) return Promise.reject(new Error('未登录'))
- return createAISession(session, projectId, input)
+ return createAISession(session, projectId, input, signal)
}, [session])
const dark = theme === 'dark'
diff --git a/apps/web_v1/src/pages/projects/project-ai.tsx b/apps/web_v1/src/pages/projects/project-ai.tsx
index 0b55a11..cdda9bd 100644
--- a/apps/web_v1/src/pages/projects/project-ai.tsx
+++ b/apps/web_v1/src/pages/projects/project-ai.tsx
@@ -1,4 +1,4 @@
-import { useEffect, useState } from 'react'
+import { useEffect, useRef, useState } from 'react'
import { Alert, Button, Card, Empty, Input, Space, Spin, Tag, Typography } from '@arco-design/web-react'
import { IconPlusCircle, IconRobot } from '@arco-design/web-react/icon'
import type { AISessionDTO, CreateAISessionInput } from '../../api/ai'
@@ -14,8 +14,8 @@ export function ProjectAi({
}: {
activeWorkspace: ProjectWorkspace
onSelectItem: (title: string) => void
- onListSessions: (projectId: string) => Promise
- onCreateSession: (projectId: string, input: CreateAISessionInput) => Promise
+ onListSessions: (projectId: string, signal?: AbortSignal) => Promise
+ onCreateSession: (projectId: string, input: CreateAISessionInput, signal?: AbortSignal) => Promise
}) {
const [sessions, setSessions] = useState([])
const [title, setTitle] = useState('')
@@ -24,23 +24,47 @@ export function ProjectAi({
const [creating, setCreating] = useState(false)
const [error, setError] = useState('')
const projectId = activeWorkspace.project.id
+ const projectRef = useRef(projectId)
+ const generationRef = useRef(0)
+ const listControllerRef = useRef(null)
+ const createControllerRef = useRef(null)
+ projectRef.current = projectId
useEffect(() => {
- let current = true
+ const generation = ++generationRef.current
+ const controller = new AbortController()
+ listControllerRef.current?.abort()
+ createControllerRef.current?.abort()
+ listControllerRef.current = controller
+ createControllerRef.current = null
+ setSessions([])
+ setTitle('')
+ setContext('')
setLoading(true)
+ setCreating(false)
setError('')
- void onListSessions(projectId)
+ const isCurrent = () => projectRef.current === projectId && generationRef.current === generation && !controller.signal.aborted
+
+ void onListSessions(projectId, controller.signal)
.then((items) => {
- if (current) setSessions(items)
+ if (isCurrent()) setSessions(items)
})
.catch((requestError: unknown) => {
- if (current) setError(requestError instanceof Error ? requestError.message : 'AI 会话加载失败,请稍后重试')
+ if (isCurrent()) setError(requestError instanceof Error ? requestError.message : 'AI 会话加载失败,请稍后重试')
})
.finally(() => {
- if (current) setLoading(false)
+ if (isCurrent()) {
+ setLoading(false)
+ listControllerRef.current = null
+ }
})
+
return () => {
- current = false
+ generationRef.current += 1
+ controller.abort()
+ if (listControllerRef.current === controller) listControllerRef.current = null
+ createControllerRef.current?.abort()
+ createControllerRef.current = null
}
}, [onListSessions, projectId])
@@ -50,21 +74,35 @@ export function ProjectAi({
setError('请输入 AI 会话标题')
return
}
+ const generation = ++generationRef.current
+ listControllerRef.current?.abort()
+ listControllerRef.current = null
+ createControllerRef.current?.abort()
+ const controller = new AbortController()
+ createControllerRef.current = controller
+ const isCurrent = () => projectRef.current === projectId && generationRef.current === generation && !controller.signal.aborted
+
+ setLoading(false)
setCreating(true)
setError('')
try {
const created = await onCreateSession(projectId, {
title: trimmedTitle,
context: context.trim(),
- })
+ }, controller.signal)
+ if (!isCurrent()) return
setSessions((current) => [created, ...current.filter((session) => session.id !== created.id)])
setTitle('')
setContext('')
onSelectItem(created.title)
} catch (requestError) {
+ if (!isCurrent()) return
setError(requestError instanceof Error ? requestError.message : 'AI 会话创建失败,请稍后重试')
} finally {
- setCreating(false)
+ if (isCurrent()) {
+ setCreating(false)
+ createControllerRef.current = null
+ }
}
}
diff --git a/apps/web_v1/src/pages/projects/project-channel-page.tsx b/apps/web_v1/src/pages/projects/project-channel-page.tsx
index 6809187..fd4f081 100644
--- a/apps/web_v1/src/pages/projects/project-channel-page.tsx
+++ b/apps/web_v1/src/pages/projects/project-channel-page.tsx
@@ -40,8 +40,8 @@ export function ProjectChannelPage({
onUpdateTask: (update: ProjectTaskUpdate) => void
onAnalyzeInbox: (inboxId: string) => Promise
onConfirmInbox: (inboxId: string, suggestionIds: string[]) => Promise
- onListAISessions: (projectId: string) => Promise
- onCreateAISession: (projectId: string, input: CreateAISessionInput) => Promise
+ onListAISessions: (projectId: string, signal?: AbortSignal) => Promise
+ onCreateAISession: (projectId: string, input: CreateAISessionInput, signal?: AbortSignal) => Promise
}) {
switch (activeChannel) {
case 'inbox':
@@ -49,7 +49,7 @@ export function ProjectChannelPage({
case 'tasks':
return
case 'ai':
- return
+ return
case 'notes':
return
case 'cron':
diff --git a/apps/web_v1/src/pages/projects/project-statusbar.tsx b/apps/web_v1/src/pages/projects/project-statusbar.tsx
index b03e286..86b1254 100644
--- a/apps/web_v1/src/pages/projects/project-statusbar.tsx
+++ b/apps/web_v1/src/pages/projects/project-statusbar.tsx
@@ -1,5 +1,4 @@
-import { Layout, Space } from '@arco-design/web-react'
-import { IconCheckCircle } from '@arco-design/web-react/icon'
+import { Layout } from '@arco-design/web-react'
const { Footer } = Layout
@@ -10,10 +9,6 @@ export function ProjectStatusbar() {
森
已登录
-
-
- 服务已连接
-
)
}
diff --git a/apps/web_v1/src/pages/workspace-home.tsx b/apps/web_v1/src/pages/workspace-home.tsx
index bcc40d8..a5877f9 100644
--- a/apps/web_v1/src/pages/workspace-home.tsx
+++ b/apps/web_v1/src/pages/workspace-home.tsx
@@ -78,8 +78,8 @@ export function ProjectPage({
onSelectSearchResult: (result: SearchResultDTO) => void
onAnalyzeInbox: (inboxId: string) => Promise
onConfirmInbox: (inboxId: string, suggestionIds: string[]) => Promise
- onListAISessions: (projectId: string) => Promise
- onCreateAISession: (projectId: string, input: CreateAISessionInput) => Promise
+ onListAISessions: (projectId: string, signal?: AbortSignal) => Promise
+ onCreateAISession: (projectId: string, input: CreateAISessionInput, signal?: AbortSignal) => Promise
}) {
const isProject = activeView === 'project'
const projects = workspaces.map((workspace) => workspace.project)
diff --git a/backend/internal/logic/ai/gateway.go b/backend/internal/logic/ai/gateway.go
index d8889e5..9fb43c2 100644
--- a/backend/internal/logic/ai/gateway.go
+++ b/backend/internal/logic/ai/gateway.go
@@ -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
diff --git a/backend/internal/logic/ai/gateway_test.go b/backend/internal/logic/ai/gateway_test.go
index b85c21e..0e88694 100644
--- a/backend/internal/logic/ai/gateway_test.go
+++ b/backend/internal/logic/ai/gateway_test.go
@@ -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")
}
diff --git a/backend/internal/logic/ai/handlers_test.go b/backend/internal/logic/ai/handlers_test.go
index fa552d7..4491e41 100644
--- a/backend/internal/logic/ai/handlers_test.go
+++ b/backend/internal/logic/ai/handlers_test.go
@@ -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
}
diff --git a/backend/internal/logic/ai/rate_limit_postgres_test.go b/backend/internal/logic/ai/rate_limit_postgres_test.go
new file mode 100644
index 0000000..9252b77
--- /dev/null
+++ b/backend/internal/logic/ai/rate_limit_postgres_test.go
@@ -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)
+}
diff --git a/backend/internal/logic/ai/rate_limit_test.go b/backend/internal/logic/ai/rate_limit_test.go
new file mode 100644
index 0000000..94db4b0
--- /dev/null
+++ b/backend/internal/logic/ai/rate_limit_test.go
@@ -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
+}
diff --git a/backend/internal/logic/ai/sessions.go b/backend/internal/logic/ai/sessions.go
index 3e59d48..3214e56 100644
--- a/backend/internal/logic/ai/sessions.go
+++ b/backend/internal/logic/ai/sessions.go
@@ -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
diff --git a/backend/internal/models/ai_rate_bucket.go b/backend/internal/models/ai_rate_bucket.go
new file mode 100644
index 0000000..d2fd234
--- /dev/null
+++ b/backend/internal/models/ai_rate_bucket.go
@@ -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"
+}
diff --git a/backend/internal/models/new.go b/backend/internal/models/new.go
index 89fe3de..f001481 100644
--- a/backend/internal/models/new.go
+++ b/backend/internal/models/new.go
@@ -38,6 +38,7 @@ func AutoMigrate(database *gorm.DB) error {
&SenlinAgentProjectEvent{},
&SenlinAgentAIKey{},
&SenlinAgentAICallLog{},
+ &SenlinAgentAIRateBucket{},
&SenlinAgentTaskShare{},
)
}