fix: harden controlled AI session consistency
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user