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

279 lines
11 KiB
Go

package ai
import (
"bytes"
"encoding/json"
"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 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.RecordCall(owner.ID, "openai", "system", "ai_session_create", "ready", ""))
}
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)
}
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", "createdAt", "updatedAt"}, aiMapKeys(payload))
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", "createdAt", "updatedAt"}, aiMapKeys(payload[0]))
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) CheckRateLimit(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(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
}