387 lines
16 KiB
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.SenlinAgentAIExpert
|
|
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
|
|
}
|