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/initdb" "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.SaAISession{}).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.SaAIExpertItem 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.SaAISession 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.SaAICallLog 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.SaAISession{}, &models.SaTask{}, &models.SaNote{}, &models.SaSource{}, } { var count int64 require.NoError(t, database.Model(model).Count(&count).Error) require.Zero(t, count) } var call models.SaAICallLog 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.SaAIRateBucket 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.SaAICallLog) 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.SaAISession{}).Where("project_id = ?", project.ID).Count(&sessionCount).Error) require.Zero(t, sessionCount) var calls []models.SaAICallLog 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.SaAIKey{ 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.SaAICallLog 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.SaAISession{}).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.SaAICallLog 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.SaTask{}, &models.SaNote{}, &models.SaSource{}} { 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.SaAISession{ ProjectID: project.ID, CreatedBy: owner.ID, Title: "目标会话", Context: "项目上下文", }).Error) require.NoError(t, database.Create(&models.SaAISession{ 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)) require.NoError(t, initdb.InitExpert(database)) models.DBService = database return database } func createAIHandlerUser(t *testing.T, database *gorm.DB, email string) models.SaUser { t.Helper() user := models.SaUser{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.SaProject { t.Helper() project := models.SaProject{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 }