Files
agent/backend/internal/logic/ai/handlers_test.go

387 lines
16 KiB
Go

package ai
import (
"bytes"
"encoding/json"
"errors"
"fmt"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/glebarez/sqlite"
"github.com/google/uuid"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
"gorm.io/gorm/logger"
"senlinai-agent/backend/internal/config"
"senlinai-agent/backend/internal/httpx"
"senlinai-agent/backend/internal/models"
)
func TestAISessionHandlersRequireOwnedProject(t *testing.T) {
database := newAIHandlerTestDB(t)
owner := createAIHandlerUser(t, database, "owner@example.com")
intruder := createAIHandlerUser(t, database, "intruder@example.com")
project := createAIHandlerProject(t, database, owner.ID, "PRIVATE")
gateway := &recordingSessionGateway{}
router := aiHandlerTestRouter(intruder.ID, gateway)
for _, request := range []*http.Request{
authenticatedAIRequest(t, http.MethodGet, "/api/v1/projects/"+project.Identity+"/ai-sessions", nil),
authenticatedAIRequest(t, http.MethodPost, "/api/v1/projects/"+project.Identity+"/ai-sessions", map[string]any{
"title": "越权会话", "context": "不得创建",
}),
} {
recorder := httptest.NewRecorder()
router.ServeHTTP(recorder, request)
require.Equal(t, http.StatusNotFound, recorder.Code)
var payload httpx.ErrorEnvelope
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &payload))
require.Equal(t, "not_found", payload.Error.Code)
require.Equal(t, "项目不存在或无权访问", payload.Error.Message)
}
var count int64
require.NoError(t, database.Model(&models.SenlinAgentAISession{}).Count(&count).Error)
require.Zero(t, count)
require.Empty(t, gateway.steps)
}
func TestAIExpertLibraryListsDetailsAndCreatesAssociatedSession(t *testing.T) {
database := newAIHandlerTestDB(t)
owner := createAIHandlerUser(t, database, "expert-owner@example.com")
project := createAIHandlerProject(t, database, owner.ID, "EXPERTS")
var expert models.SenlinAgentAIExpertItem
require.NoError(t, database.Order("id asc").First(&expert).Error)
router := aiHandlerTestRouter(owner.ID, NewGatewayWithSecret("system-key", "test-encryption-secret"))
listRecorder := httptest.NewRecorder()
router.ServeHTTP(listRecorder, authenticatedAIRequest(t, http.MethodGet, "/api/v1/ai-experts", nil))
require.Equal(t, http.StatusOK, listRecorder.Code)
var listPayload []map[string]any
require.NoError(t, json.Unmarshal(listRecorder.Body.Bytes(), &listPayload))
require.Len(t, listPayload, 267)
require.Equal(t, expert.Identity, listPayload[0]["id"])
require.NotContains(t, listPayload[0], "systemPrompt")
detailRecorder := httptest.NewRecorder()
router.ServeHTTP(detailRecorder, authenticatedAIRequest(t, http.MethodGet, "/api/v1/ai-experts/"+expert.Identity, nil))
require.Equal(t, http.StatusOK, detailRecorder.Code)
var detailPayload map[string]any
require.NoError(t, json.Unmarshal(detailRecorder.Body.Bytes(), &detailPayload))
require.Equal(t, expert.Name, detailPayload["name"])
require.NotEmpty(t, detailPayload["systemPrompt"])
require.Equal(t, "MIT", detailPayload["sourceLicense"])
createRecorder := httptest.NewRecorder()
router.ServeHTTP(createRecorder, authenticatedAIRequest(t, http.MethodPost, "/api/v1/projects/"+project.Identity+"/ai-sessions", map[string]any{
"title": "专家会话", "context": "请给出下一步计划", "expertId": expert.Identity,
}))
require.Equal(t, http.StatusCreated, createRecorder.Code)
var createPayload map[string]any
require.NoError(t, json.Unmarshal(createRecorder.Body.Bytes(), &createPayload))
responseExpert := createPayload["expert"].(map[string]any)
require.Equal(t, expert.Identity, responseExpert["id"])
require.Equal(t, expert.Name, responseExpert["name"])
var session models.SenlinAgentAISession
require.NoError(t, database.First(&session).Error)
require.NotNil(t, session.ExpertID)
require.Equal(t, expert.ID, *session.ExpertID)
}
func TestCreateAISessionChecksRateLimitBeforeSelectingProvider(t *testing.T) {
database := newAIHandlerTestDB(t)
owner := createAIHandlerUser(t, database, "owner@example.com")
project := createAIHandlerProject(t, database, owner.ID, "RATE")
gateway := &recordingSessionGateway{
selected: SelectedKey{Provider: "openai", APIKey: "system-key", KeyType: "system"},
}
router := aiHandlerTestRouter(owner.ID, gateway)
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.StatusCreated, recorder.Code)
require.Equal(t, []string{"rate", "select", "record"}, gateway.steps)
}
func TestCreateAISessionReturnsRateLimitBeforeMissingKeyAndAuditsFailure(t *testing.T) {
database := newAIHandlerTestDB(t)
owner := createAIHandlerUser(t, database, "owner@example.com")
project := createAIHandlerProject(t, database, owner.ID, "LIMITED")
gateway := NewGatewayWithSecret("", "test-encryption-secret")
for range aiSessionCreateLimit {
require.NoError(t, gateway.ReserveRateLimit(owner.ID, aiSessionCreateAction, aiSessionCreateLimit, time.Hour))
}
router := aiHandlerTestRouter(owner.ID, gateway)
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.StatusTooManyRequests, recorder.Code)
var payload httpx.ErrorEnvelope
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &payload))
require.Equal(t, "ai_rate_limited", payload.Error.Code)
require.Equal(t, "AI 请求过于频繁,请稍后重试", payload.Error.Message)
var latest models.SenlinAgentAICallLog
require.NoError(t, database.Order("id desc").First(&latest).Error)
require.Equal(t, "none", latest.Provider)
require.Equal(t, "none", latest.UsedKeyType)
require.Equal(t, "ai_session_create", latest.Action)
require.Equal(t, "failed", latest.Status)
require.Equal(t, "ai_rate_limited", latest.Error)
}
func TestCreateAISessionWithoutKeyReturnsAuditedErrorAndCreatesNoFormalObjects(t *testing.T) {
database := newAIHandlerTestDB(t)
owner := createAIHandlerUser(t, database, "owner@example.com")
project := createAIHandlerProject(t, database, owner.ID, "NO_KEY")
router := aiHandlerTestRouter(owner.ID, NewGatewayWithSecret("", "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.StatusServiceUnavailable, recorder.Code)
var payload httpx.ErrorEnvelope
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &payload))
require.Equal(t, "ai_key_missing", payload.Error.Code)
require.Equal(t, "尚未配置可用的 AI 密钥", payload.Error.Message)
for _, model := range []any{
&models.SenlinAgentAISession{},
&models.SenlinAgentTask{},
&models.SenlinAgentNote{},
&models.SenlinAgentSource{},
} {
var count int64
require.NoError(t, database.Model(model).Count(&count).Error)
require.Zero(t, count)
}
var call models.SenlinAgentAICallLog
require.NoError(t, database.First(&call).Error)
require.Equal(t, owner.ID, call.UserID)
require.Equal(t, "none", call.Provider)
require.Equal(t, "none", call.UsedKeyType)
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) {
database := newAIHandlerTestDB(t)
owner := createAIHandlerUser(t, database, "owner@example.com")
project := createAIHandlerProject(t, database, owner.ID, "CREATE")
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.StatusCreated, recorder.Code)
var payload map[string]any
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &payload))
require.ElementsMatch(t, []string{"id", "projectId", "title", "context", "status", "expert", "createdAt", "updatedAt"}, aiMapKeys(payload))
require.Nil(t, payload["expert"])
identity, err := uuid.Parse(payload["id"].(string))
require.NoError(t, err)
require.Equal(t, uuid.Version(7), identity.Version())
require.Equal(t, project.Identity, payload["projectId"])
require.Equal(t, "报价分析", payload["title"])
require.Equal(t, "仅整理会话上下文", payload["context"])
require.Equal(t, "ready", payload["status"])
for _, forbidden := range []string{"taskId", "noteId", "sourceId", "createdTaskId", "createdNoteId", "createdSourceId"} {
require.NotContains(t, payload, forbidden)
}
var call models.SenlinAgentAICallLog
require.NoError(t, database.First(&call).Error)
require.Equal(t, "openai", call.Provider)
require.Equal(t, "system", call.UsedKeyType)
require.Equal(t, "ai_session_create", call.Action)
require.Equal(t, "ready", call.Status)
require.Empty(t, call.Error)
for _, model := range []any{&models.SenlinAgentTask{}, &models.SenlinAgentNote{}, &models.SenlinAgentSource{}} {
var count int64
require.NoError(t, database.Model(model).Count(&count).Error)
require.Zero(t, count)
}
}
func TestListAISessionsReturnsOnlyOwnedProjectIdentityDTOs(t *testing.T) {
database := newAIHandlerTestDB(t)
owner := createAIHandlerUser(t, database, "owner@example.com")
project := createAIHandlerProject(t, database, owner.ID, "LIST")
otherProject := createAIHandlerProject(t, database, owner.ID, "OTHER")
require.NoError(t, database.Create(&models.SenlinAgentAISession{
ProjectID: project.ID, CreatedBy: owner.ID, Title: "目标会话", Context: "项目上下文",
}).Error)
require.NoError(t, database.Create(&models.SenlinAgentAISession{
ProjectID: otherProject.ID, CreatedBy: owner.ID, Title: "其他会话", Context: "不得混入",
}).Error)
router := aiHandlerTestRouter(owner.ID, &recordingSessionGateway{})
recorder := httptest.NewRecorder()
router.ServeHTTP(recorder, authenticatedAIRequest(t, http.MethodGet, "/api/v1/projects/"+project.Identity+"/ai-sessions", nil))
require.Equal(t, http.StatusOK, recorder.Code)
var payload []map[string]any
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &payload))
require.Len(t, payload, 1)
require.ElementsMatch(t, []string{"id", "projectId", "title", "context", "status", "expert", "createdAt", "updatedAt"}, aiMapKeys(payload[0]))
require.Nil(t, payload[0]["expert"])
require.Equal(t, project.Identity, payload[0]["projectId"])
require.Equal(t, "目标会话", payload[0]["title"])
require.Equal(t, "ready", payload[0]["status"])
}
type recordingSessionGateway struct {
steps []string
selected SelectedKey
rateErr error
selectErr error
}
func (g *recordingSessionGateway) ReserveRateLimit(uint, string, int, time.Duration) error {
g.steps = append(g.steps, "rate")
return g.rateErr
}
func (g *recordingSessionGateway) SelectKey(uint) (SelectedKey, error) {
g.steps = append(g.steps, "select")
return g.selected, g.selectErr
}
func (g *recordingSessionGateway) RecordCall(*gorm.DB, uint, string, string, string, string, string) error {
g.steps = append(g.steps, "record")
return nil
}
func newAIHandlerTestDB(t *testing.T) *gorm.DB {
t.Helper()
database, err := gorm.Open(sqlite.Open(fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
require.NoError(t, err)
require.NoError(t, models.AutoMigrate(database))
models.DBService = database
return database
}
func createAIHandlerUser(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
}
func createAIHandlerProject(t *testing.T, database *gorm.DB, ownerID uint, identifier string) models.SenlinAgentProject {
t.Helper()
project := models.SenlinAgentProject{OwnerID: ownerID, Name: identifier, Identifier: identifier}
require.NoError(t, database.Create(&project).Error)
return project
}
func aiHandlerTestRouter(userID uint, gateway sessionGateway) http.Handler {
return httpx.NewProtectedRouter(
config.Config{Env: "test"},
func(string) (uint, error) { return userID, nil },
NewHandler(NewSessionService(gateway)),
)
}
func authenticatedAIRequest(t *testing.T, method, path string, body any) *http.Request {
t.Helper()
var requestBody *bytes.Reader
if body == nil {
requestBody = bytes.NewReader(nil)
} else {
encoded, err := json.Marshal(body)
require.NoError(t, err)
requestBody = bytes.NewReader(encoded)
}
request := httptest.NewRequest(method, path, requestBody)
request.Header.Set("Authorization", "Bearer test-token")
if body != nil {
request.Header.Set("Content-Type", "application/json")
}
return request
}
func aiMapKeys(values map[string]any) []string {
keys := make([]string, 0, len(values))
for key := range values {
keys = append(keys, key)
}
return keys
}