diff --git a/backend/internal/api/conversation_test.go b/backend/internal/api/conversation_test.go new file mode 100644 index 0000000..6b873fc --- /dev/null +++ b/backend/internal/api/conversation_test.go @@ -0,0 +1,575 @@ +package api_test + +import ( + "bytes" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/hhs/camtalk/internal/api" + "github.com/hhs/camtalk/internal/auth" + "github.com/hhs/camtalk/internal/models" + "github.com/hhs/camtalk/internal/session" +) + +// mockSessionManager 实现 session.Manager 接口,用于 ConversationHandler 测试。 +type mockSessionManager struct { + CreateFunc func(ctx context.Context, userID string, config models.SessionConfig) (string, error) + GetFunc func(ctx context.Context, sessionID string) (*models.Session, error) + UpdateConfigFunc func(ctx context.Context, sessionID string, patch models.SessionConfigPatch) error + UpdateTitleFunc func(ctx context.Context, sessionID string, title string) error + ListByUserFunc func(ctx context.Context, userID string, page, size int) ([]session.ConversationSummary, int, error) + GetHistoryFunc func(ctx context.Context, sessionID string, limit int) ([]models.Message, error) + AppendMessageFunc func(ctx context.Context, sessionID string, msg models.Message) error + SetActiveRequestFunc func(ctx context.Context, sessionID string, requestID string) error + GetActiveRequestIDFunc func(ctx context.Context, sessionID string) (string, error) + ClearActiveRequestFunc func(ctx context.Context, sessionID string) error + TouchFunc func(ctx context.Context, sessionID string) error + DestroyFunc func(ctx context.Context, sessionID string) error + ActiveCountFunc func() int +} + +func (m *mockSessionManager) Create(ctx context.Context, userID string, config models.SessionConfig) (string, error) { + return m.CreateFunc(ctx, userID, config) +} + +func (m *mockSessionManager) Get(ctx context.Context, sessionID string) (*models.Session, error) { + return m.GetFunc(ctx, sessionID) +} + +func (m *mockSessionManager) UpdateConfig(ctx context.Context, sessionID string, patch models.SessionConfigPatch) error { + return m.UpdateConfigFunc(ctx, sessionID, patch) +} + +func (m *mockSessionManager) UpdateTitle(ctx context.Context, sessionID string, title string) error { + return m.UpdateTitleFunc(ctx, sessionID, title) +} + +func (m *mockSessionManager) ListByUser(ctx context.Context, userID string, page, size int) ([]session.ConversationSummary, int, error) { + return m.ListByUserFunc(ctx, userID, page, size) +} + +func (m *mockSessionManager) GetHistory(ctx context.Context, sessionID string, limit int) ([]models.Message, error) { + return m.GetHistoryFunc(ctx, sessionID, limit) +} + +func (m *mockSessionManager) AppendMessage(ctx context.Context, sessionID string, msg models.Message) error { + return m.AppendMessageFunc(ctx, sessionID, msg) +} + +func (m *mockSessionManager) SetActiveRequest(ctx context.Context, sessionID string, requestID string) error { + return m.SetActiveRequestFunc(ctx, sessionID, requestID) +} + +func (m *mockSessionManager) GetActiveRequestID(ctx context.Context, sessionID string) (string, error) { + return m.GetActiveRequestIDFunc(ctx, sessionID) +} + +func (m *mockSessionManager) ClearActiveRequest(ctx context.Context, sessionID string) error { + return m.ClearActiveRequestFunc(ctx, sessionID) +} + +func (m *mockSessionManager) Touch(ctx context.Context, sessionID string) error { + return m.TouchFunc(ctx, sessionID) +} + +func (m *mockSessionManager) Destroy(ctx context.Context, sessionID string) error { + return m.DestroyFunc(ctx, sessionID) +} + +func (m *mockSessionManager) ActiveCount() int { + return m.ActiveCountFunc() +} + +// newConvTestRouter 创建带 ConversationHandler 路由的测试引擎,同时返回 TokenManager。 +func newConvTestRouter(mgr session.Manager) (*gin.Engine, *auth.TokenManager) { + gin.SetMode(gin.TestMode) + r := gin.New() + tm := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour) + h := api.NewConversationHandler(mgr, tm) + h.RegisterRoutes(r.Group("/api")) + return r, tm +} + +// --- List --- + +func TestConversationList_Success(t *testing.T) { + now := time.Now() + mgr := &mockSessionManager{ + ListByUserFunc: func(_ context.Context, userID string, page, size int) ([]session.ConversationSummary, int, error) { + assert.Equal(t, "user-123", userID) + assert.Equal(t, 1, page) + assert.Equal(t, 20, size) + return []session.ConversationSummary{ + {ID: "sess-1", Title: "对话一", MessageCount: 3, UpdatedAt: now}, + {ID: "sess-2", Title: "对话二", MessageCount: 1, UpdatedAt: now.Add(-time.Hour)}, + }, 2, nil + }, + } + r, tm := newConvTestRouter(mgr) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + req := httptest.NewRequest(http.MethodGet, "/api/conversations", nil) + req.Header.Set("Authorization", "Bearer "+access) + w := httptest.NewRecorder() + + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusOK, w.Code) + var resp map[string]interface{} + require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp)) + assert.Equal(t, float64(2), resp["total"]) + convs := resp["conversations"].([]interface{}) + assert.Len(t, convs, 2) +} + +func TestConversationList_WithPagination(t *testing.T) { + mgr := &mockSessionManager{ + ListByUserFunc: func(_ context.Context, _ string, page, size int) ([]session.ConversationSummary, int, error) { + assert.Equal(t, 2, page) + assert.Equal(t, 10, size) + return []session.ConversationSummary{}, 0, nil + }, + } + r, tm := newConvTestRouter(mgr) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + req := httptest.NewRequest(http.MethodGet, "/api/conversations?page=2&size=10", nil) + req.Header.Set("Authorization", "Bearer "+access) + w := httptest.NewRecorder() + + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusOK, w.Code) +} + +func TestConversationList_MissingAuth(t *testing.T) { + mgr := &mockSessionManager{} + r, _ := newConvTestRouter(mgr) + + req := httptest.NewRequest(http.MethodGet, "/api/conversations", nil) + w := httptest.NewRecorder() + + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusUnauthorized, w.Code) +} + +// --- Create --- + +func TestConversationCreate_Success(t *testing.T) { + createdID := "new-session-id" + now := time.Now() + mgr := &mockSessionManager{ + CreateFunc: func(_ context.Context, userID string, cfg models.SessionConfig) (string, error) { + assert.Equal(t, "user-123", userID) + return createdID, nil + }, + GetFunc: func(_ context.Context, sessionID string) (*models.Session, error) { + assert.Equal(t, createdID, sessionID) + return &models.Session{ + ID: createdID, + UserID: "user-123", + Title: models.DefaultSessionTitle, + CreatedAt: now, + UpdatedAt: now, + Config: models.DefaultConfig(), + }, nil + }, + } + r, tm := newConvTestRouter(mgr) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + req := httptest.NewRequest(http.MethodPost, "/api/conversations", nil) + req.Header.Set("Authorization", "Bearer "+access) + w := httptest.NewRecorder() + + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusCreated, w.Code) + var resp map[string]interface{} + require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp)) + assert.Equal(t, createdID, resp["id"]) + assert.Equal(t, models.DefaultSessionTitle, resp["title"]) +} + +func TestConversationCreate_WithConfig(t *testing.T) { + mgr := &mockSessionManager{ + CreateFunc: func(_ context.Context, _ string, cfg models.SessionConfig) (string, error) { + assert.False(t, cfg.TTSEnabled) + assert.Equal(t, "high", cfg.DetailLevel) + return "sess-1", nil + }, + GetFunc: func(_ context.Context, _ string) (*models.Session, error) { + return &models.Session{ + ID: "sess-1", + UserID: "user-123", + Title: models.DefaultSessionTitle, + CreatedAt: time.Now(), + UpdatedAt: time.Now(), + }, nil + }, + } + r, tm := newConvTestRouter(mgr) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + body, _ := json.Marshal(api.CreateConversationRequest{ + Config: &models.SessionConfig{TTSEnabled: false, DetailLevel: "high", Language: "zh-CN"}, + }) + req := httptest.NewRequest(http.MethodPost, "/api/conversations", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+access) + w := httptest.NewRecorder() + + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusCreated, w.Code) +} + +// --- Get --- + +func TestConversationGet_Success(t *testing.T) { + now := time.Now() + mgr := &mockSessionManager{ + GetFunc: func(_ context.Context, sessionID string) (*models.Session, error) { + assert.Equal(t, "sess-1", sessionID) + return &models.Session{ + ID: "sess-1", + UserID: "user-123", + Title: "我的对话", + CreatedAt: now, + UpdatedAt: now, + Config: models.DefaultConfig(), + }, nil + }, + } + r, tm := newConvTestRouter(mgr) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + req := httptest.NewRequest(http.MethodGet, "/api/conversations/sess-1", nil) + req.Header.Set("Authorization", "Bearer "+access) + w := httptest.NewRecorder() + + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusOK, w.Code) + var resp map[string]interface{} + require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp)) + assert.Equal(t, "我的对话", resp["title"]) +} + +func TestConversationGet_NotFound(t *testing.T) { + mgr := &mockSessionManager{ + GetFunc: func(_ context.Context, _ string) (*models.Session, error) { + return nil, session.ErrSessionNotFound + }, + } + r, tm := newConvTestRouter(mgr) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + req := httptest.NewRequest(http.MethodGet, "/api/conversations/nonexistent", nil) + req.Header.Set("Authorization", "Bearer "+access) + w := httptest.NewRecorder() + + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusNotFound, w.Code) + assert.Contains(t, w.Body.String(), "SESSION_NOT_FOUND") +} + +func TestConversationGet_Forbidden(t *testing.T) { + mgr := &mockSessionManager{ + GetFunc: func(_ context.Context, _ string) (*models.Session, error) { + // 会话属于另一个用户 + return &models.Session{ + ID: "sess-1", + UserID: "other-user", + Title: "他人对话", + }, nil + }, + } + r, tm := newConvTestRouter(mgr) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + req := httptest.NewRequest(http.MethodGet, "/api/conversations/sess-1", nil) + req.Header.Set("Authorization", "Bearer "+access) + w := httptest.NewRecorder() + + r.ServeHTTP(w, req) + + // 返回 404 而非 403,避免信息泄露 + assert.Equal(t, http.StatusNotFound, w.Code) + assert.Contains(t, w.Body.String(), "SESSION_NOT_FOUND") +} + +// --- UpdateTitle --- + +func TestConversationUpdateTitle_Success(t *testing.T) { + mgr := &mockSessionManager{ + GetFunc: func(_ context.Context, _ string) (*models.Session, error) { + return &models.Session{ID: "sess-1", UserID: "user-123"}, nil + }, + UpdateTitleFunc: func(_ context.Context, sessionID, title string) error { + assert.Equal(t, "sess-1", sessionID) + assert.Equal(t, "新标题", title) + return nil + }, + } + r, tm := newConvTestRouter(mgr) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + body, _ := json.Marshal(api.UpdateTitleRequest{Title: "新标题"}) + req := httptest.NewRequest(http.MethodPatch, "/api/conversations/sess-1", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+access) + w := httptest.NewRecorder() + + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusOK, w.Code) + assert.Contains(t, w.Body.String(), "title updated") +} + +func TestConversationUpdateTitle_EmptyTitle(t *testing.T) { + mgr := &mockSessionManager{ + GetFunc: func(_ context.Context, _ string) (*models.Session, error) { + return &models.Session{ID: "sess-1", UserID: "user-123"}, nil + }, + } + r, tm := newConvTestRouter(mgr) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + body, _ := json.Marshal(api.UpdateTitleRequest{Title: ""}) + req := httptest.NewRequest(http.MethodPatch, "/api/conversations/sess-1", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+access) + w := httptest.NewRecorder() + + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusBadRequest, w.Code) + assert.Contains(t, w.Body.String(), "title is required") +} + +func TestConversationUpdateTitle_TooLong(t *testing.T) { + mgr := &mockSessionManager{ + GetFunc: func(_ context.Context, _ string) (*models.Session, error) { + return &models.Session{ID: "sess-1", UserID: "user-123"}, nil + }, + } + r, tm := newConvTestRouter(mgr) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + longTitle := "" + for i := 0; i < 101; i++ { + longTitle += "测" + } + body, _ := json.Marshal(api.UpdateTitleRequest{Title: longTitle}) + req := httptest.NewRequest(http.MethodPatch, "/api/conversations/sess-1", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+access) + w := httptest.NewRecorder() + + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusBadRequest, w.Code) + assert.Contains(t, w.Body.String(), "title must be 100 characters or less") +} + +// --- Delete --- + +func TestConversationDelete_Success(t *testing.T) { + destroyCalled := false + mgr := &mockSessionManager{ + GetFunc: func(_ context.Context, _ string) (*models.Session, error) { + return &models.Session{ID: "sess-1", UserID: "user-123"}, nil + }, + DestroyFunc: func(_ context.Context, sessionID string) error { + assert.Equal(t, "sess-1", sessionID) + destroyCalled = true + return nil + }, + } + r, tm := newConvTestRouter(mgr) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + req := httptest.NewRequest(http.MethodDelete, "/api/conversations/sess-1", nil) + req.Header.Set("Authorization", "Bearer "+access) + w := httptest.NewRecorder() + + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusNoContent, w.Code) + assert.True(t, destroyCalled) +} + +func TestConversationDelete_Forbidden(t *testing.T) { + mgr := &mockSessionManager{ + GetFunc: func(_ context.Context, _ string) (*models.Session, error) { + return &models.Session{ID: "sess-1", UserID: "other-user"}, nil + }, + } + r, tm := newConvTestRouter(mgr) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + req := httptest.NewRequest(http.MethodDelete, "/api/conversations/sess-1", nil) + req.Header.Set("Authorization", "Bearer "+access) + w := httptest.NewRecorder() + + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusNotFound, w.Code) +} + +// --- GetMessages --- + +func TestConversationGetMessages_Success(t *testing.T) { + mgr := &mockSessionManager{ + GetFunc: func(_ context.Context, _ string) (*models.Session, error) { + return &models.Session{ID: "sess-1", UserID: "user-123"}, nil + }, + GetHistoryFunc: func(_ context.Context, sessionID string, limit int) ([]models.Message, error) { + assert.Equal(t, "sess-1", sessionID) + assert.Equal(t, 0, limit) // 获取全量 + return []models.Message{ + {Role: "user", Content: "你好"}, + {Role: "assistant", Content: "你好!有什么可以帮助你的吗?"}, + {Role: "user", Content: "今天天气怎么样?"}, + {Role: "assistant", Content: "今天天气不错!"}, + }, nil + }, + } + r, tm := newConvTestRouter(mgr) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + req := httptest.NewRequest(http.MethodGet, "/api/conversations/sess-1/messages", nil) + req.Header.Set("Authorization", "Bearer "+access) + w := httptest.NewRecorder() + + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusOK, w.Code) + var resp map[string]interface{} + require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp)) + assert.Equal(t, float64(4), resp["total"]) + msgs := resp["messages"].([]interface{}) + assert.Len(t, msgs, 4) +} + +func TestConversationGetMessages_WithLimit(t *testing.T) { + mgr := &mockSessionManager{ + GetFunc: func(_ context.Context, _ string) (*models.Session, error) { + return &models.Session{ID: "sess-1", UserID: "user-123"}, nil + }, + GetHistoryFunc: func(_ context.Context, _ string, _ int) ([]models.Message, error) { + return []models.Message{ + {Role: "user", Content: "消息1"}, + {Role: "assistant", Content: "回复1"}, + {Role: "user", Content: "消息2"}, + {Role: "assistant", Content: "回复2"}, + }, nil + }, + } + r, tm := newConvTestRouter(mgr) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + req := httptest.NewRequest(http.MethodGet, "/api/conversations/sess-1/messages?limit=2", nil) + req.Header.Set("Authorization", "Bearer "+access) + w := httptest.NewRecorder() + + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusOK, w.Code) + var resp map[string]interface{} + require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp)) + msgs := resp["messages"].([]interface{}) + assert.Len(t, msgs, 2) +} + +func TestConversationGetMessages_WithBefore(t *testing.T) { + mgr := &mockSessionManager{ + GetFunc: func(_ context.Context, _ string) (*models.Session, error) { + return &models.Session{ID: "sess-1", UserID: "user-123"}, nil + }, + GetHistoryFunc: func(_ context.Context, _ string, _ int) ([]models.Message, error) { + return []models.Message{ + {Role: "user", Content: "消息1"}, + {Role: "assistant", Content: "回复1"}, + {Role: "user", Content: "消息2"}, + {Role: "assistant", Content: "回复2"}, + }, nil + }, + } + r, tm := newConvTestRouter(mgr) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + req := httptest.NewRequest(http.MethodGet, "/api/conversations/sess-1/messages?before=2&limit=10", nil) + req.Header.Set("Authorization", "Bearer "+access) + w := httptest.NewRecorder() + + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusOK, w.Code) + var resp map[string]interface{} + require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp)) + // before=2 表示取 index 0..1,共 2 条 + msgs := resp["messages"].([]interface{}) + assert.Len(t, msgs, 2) +} + +func TestConversationGetMessages_Forbidden(t *testing.T) { + mgr := &mockSessionManager{ + GetFunc: func(_ context.Context, _ string) (*models.Session, error) { + return &models.Session{ID: "sess-1", UserID: "other-user"}, nil + }, + } + r, tm := newConvTestRouter(mgr) + + access, _, err := tm.GeneratePair("user-123", "alice") + require.NoError(t, err) + + req := httptest.NewRequest(http.MethodGet, "/api/conversations/sess-1/messages", nil) + req.Header.Set("Authorization", "Bearer "+access) + w := httptest.NewRecorder() + + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusNotFound, w.Code) +}