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

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