feat: 构建用户模块,实现用户对话历史持久化,完善接口文档 #96
@@ -12,6 +12,7 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/hhs/camtalk/internal/api"
|
||||
"github.com/hhs/camtalk/internal/auth"
|
||||
"github.com/hhs/camtalk/internal/ai/llm"
|
||||
"github.com/hhs/camtalk/internal/ai/stt"
|
||||
"github.com/hhs/camtalk/internal/ai/tts"
|
||||
@@ -60,6 +61,10 @@ func main() {
|
||||
_ = pool
|
||||
}
|
||||
|
||||
// 初始化 UserRepository(内存模式用于无 DB 场景)
|
||||
var userRepo store.UserRepository
|
||||
userRepo = store.NewMemUserRepository()
|
||||
|
||||
// 初始化 Session Manager(MVP 默认内存实现)
|
||||
var sessionMgr session.Manager
|
||||
// TODO: 当 Redis 配置非空时切换为 RedisManager
|
||||
@@ -105,6 +110,14 @@ func main() {
|
||||
// 初始化 Orchestrator
|
||||
orch := orchestrator.New(sttService, llmService, ttsService, sessionMgr, cfg)
|
||||
|
||||
// 初始化认证服务
|
||||
tokenMgr := auth.NewTokenManager(
|
||||
cfg.Auth.JWTSecret,
|
||||
time.Duration(cfg.Auth.AccessTTL)*time.Minute,
|
||||
time.Duration(cfg.Auth.RefreshTTL)*time.Minute,
|
||||
)
|
||||
authService := auth.NewAuthService(tokenMgr, userRepo)
|
||||
|
||||
// Gin 模式
|
||||
if cfg.App.Env == "prod" {
|
||||
gin.SetMode(gin.ReleaseMode)
|
||||
@@ -123,8 +136,16 @@ func main() {
|
||||
sessionHandler := api.NewSessionHandler(sessionMgr)
|
||||
sessionHandler.RegisterRoutes(apiGroup)
|
||||
|
||||
// Auth REST 端点
|
||||
authHandler := api.NewAuthHandler(authService, tokenMgr)
|
||||
authHandler.RegisterRoutes(apiGroup)
|
||||
|
||||
// Conversation REST 端点
|
||||
convHandler := api.NewConversationHandler(sessionMgr, tokenMgr)
|
||||
convHandler.RegisterRoutes(apiGroup)
|
||||
|
||||
// WebSocket
|
||||
r.GET("/ws", ws.ServeWS(sessionMgr, orch, cfg))
|
||||
r.GET("/ws", ws.ServeWS(sessionMgr, orch, cfg, tokenMgr))
|
||||
|
||||
// HTTP Server
|
||||
srv := &http.Server{
|
||||
|
||||
194
backend/internal/api/auth.go
Normal file
194
backend/internal/api/auth.go
Normal file
@@ -0,0 +1,194 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/hhs/camtalk/internal/auth"
|
||||
apperr "github.com/hhs/camtalk/internal/errors"
|
||||
)
|
||||
|
||||
// AuthHandler 提供认证相关的 REST 端点。
|
||||
type AuthHandler struct {
|
||||
authService auth.Service
|
||||
tokenMgr *auth.TokenManager
|
||||
}
|
||||
|
||||
// NewAuthHandler 创建 AuthHandler。
|
||||
func NewAuthHandler(authService auth.Service, tokenMgr *auth.TokenManager) *AuthHandler {
|
||||
return &AuthHandler{
|
||||
authService: authService,
|
||||
tokenMgr: tokenMgr,
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterRoutes 注册认证相关路由到给定的路由组。
|
||||
func (h *AuthHandler) RegisterRoutes(rg *gin.RouterGroup) {
|
||||
authGroup := rg.Group("/auth")
|
||||
{
|
||||
authGroup.POST("/register", h.Register)
|
||||
authGroup.POST("/login", h.Login)
|
||||
authGroup.POST("/refresh", h.Refresh)
|
||||
authGroup.POST("/logout", auth.AuthMiddleware(h.tokenMgr), h.Logout)
|
||||
}
|
||||
}
|
||||
|
||||
// Register POST /api/auth/register — 用户注册。
|
||||
func (h *AuthHandler) Register(c *gin.Context) {
|
||||
var req auth.RegisterRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
"code": apperr.CodeInvalidInput,
|
||||
"message": "invalid request body",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if msg := validateCredentials(req.Username, req.Password); msg != "" {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
"code": apperr.CodeInvalidInput,
|
||||
"message": msg,
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
resp, err := h.authService.Register(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
handleAuthError(c, err)
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusCreated, resp)
|
||||
}
|
||||
|
||||
// Login POST /api/auth/login — 用户登录。
|
||||
func (h *AuthHandler) Login(c *gin.Context) {
|
||||
var req auth.LoginRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
"code": apperr.CodeInvalidInput,
|
||||
"message": "invalid request body",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if msg := validateCredentials(req.Username, req.Password); msg != "" {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
"code": apperr.CodeInvalidInput,
|
||||
"message": msg,
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
resp, err := h.authService.Login(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
handleAuthError(c, err)
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, resp)
|
||||
}
|
||||
|
||||
// Refresh POST /api/auth/refresh — 刷新令牌。
|
||||
func (h *AuthHandler) Refresh(c *gin.Context) {
|
||||
var req auth.RefreshRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
"code": apperr.CodeInvalidInput,
|
||||
"message": "invalid request body",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if req.RefreshToken == "" {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
"code": apperr.CodeInvalidInput,
|
||||
"message": "refresh_token is required",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
resp, err := h.authService.Refresh(c.Request.Context(), req)
|
||||
if err != nil {
|
||||
handleAuthError(c, err)
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, resp)
|
||||
}
|
||||
|
||||
// Logout POST /api/auth/logout — 登出(需要认证)。
|
||||
func (h *AuthHandler) Logout(c *gin.Context) {
|
||||
userID := c.GetString(auth.ContextKeyUserID)
|
||||
|
||||
var req struct {
|
||||
RefreshToken string `json:"refresh_token"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
"code": apperr.CodeInvalidInput,
|
||||
"message": "invalid request body",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if req.RefreshToken == "" {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
"code": apperr.CodeInvalidInput,
|
||||
"message": "refresh_token is required",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.authService.Logout(c.Request.Context(), userID, req.RefreshToken); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": apperr.CodeInternalError,
|
||||
"message": "failed to logout",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"message": "logged out successfully",
|
||||
})
|
||||
}
|
||||
|
||||
// validateCredentials 校验用户名和密码格式。
|
||||
// 返回空字符串表示校验通过,否则返回错误描述。
|
||||
func validateCredentials(username, password string) string {
|
||||
if len(username) < 3 || len(username) > 64 {
|
||||
return "username must be 3-64 characters"
|
||||
}
|
||||
if len(password) < 8 || len(password) > 72 {
|
||||
return "password must be 8-72 characters"
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// handleAuthError 将 auth 层错误映射为 HTTP 响应。
|
||||
func handleAuthError(c *gin.Context, err error) {
|
||||
switch {
|
||||
case errors.Is(err, auth.ErrUsernameTaken):
|
||||
c.JSON(http.StatusConflict, gin.H{
|
||||
"code": apperr.CodeUsernameTaken,
|
||||
"message": "username already taken",
|
||||
})
|
||||
case errors.Is(err, auth.ErrInvalidCredentials):
|
||||
c.JSON(http.StatusUnauthorized, gin.H{
|
||||
"code": apperr.CodeInvalidCredentials,
|
||||
"message": "invalid username or password",
|
||||
})
|
||||
case errors.Is(err, auth.ErrRefreshTokenUsed):
|
||||
c.JSON(http.StatusUnauthorized, gin.H{
|
||||
"code": apperr.CodeInvalidToken,
|
||||
"message": "refresh token has been used or expired",
|
||||
})
|
||||
default:
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": apperr.CodeInternalError,
|
||||
"message": "internal server error",
|
||||
})
|
||||
}
|
||||
}
|
||||
324
backend/internal/api/auth_test.go
Normal file
324
backend/internal/api/auth_test.go
Normal file
@@ -0,0 +1,324 @@
|
||||
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"
|
||||
)
|
||||
|
||||
// mockAuthService 实现 auth.Service 接口,用于 API 测试。
|
||||
type mockAuthService struct {
|
||||
RegisterFunc func(ctx context.Context, req auth.RegisterRequest) (*auth.AuthResponse, error)
|
||||
LoginFunc func(ctx context.Context, req auth.LoginRequest) (*auth.AuthResponse, error)
|
||||
RefreshFunc func(ctx context.Context, req auth.RefreshRequest) (*auth.AuthResponse, error)
|
||||
LogoutFunc func(ctx context.Context, userID, refreshToken string) error
|
||||
}
|
||||
|
||||
func (m *mockAuthService) Register(ctx context.Context, req auth.RegisterRequest) (*auth.AuthResponse, error) {
|
||||
return m.RegisterFunc(ctx, req)
|
||||
}
|
||||
|
||||
func (m *mockAuthService) Login(ctx context.Context, req auth.LoginRequest) (*auth.AuthResponse, error) {
|
||||
return m.LoginFunc(ctx, req)
|
||||
}
|
||||
|
||||
func (m *mockAuthService) Refresh(ctx context.Context, req auth.RefreshRequest) (*auth.AuthResponse, error) {
|
||||
return m.RefreshFunc(ctx, req)
|
||||
}
|
||||
|
||||
func (m *mockAuthService) Logout(ctx context.Context, userID, refreshToken string) error {
|
||||
return m.LogoutFunc(ctx, userID, refreshToken)
|
||||
}
|
||||
|
||||
// newTestRouter 创建带 AuthHandler 路由的测试 Gin 引擎。
|
||||
func newTestRouter(svc auth.Service) *gin.Engine {
|
||||
gin.SetMode(gin.TestMode)
|
||||
r := gin.New()
|
||||
tm := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour)
|
||||
h := api.NewAuthHandler(svc, tm)
|
||||
h.RegisterRoutes(r.Group("/api"))
|
||||
return r
|
||||
}
|
||||
|
||||
// newTestRouterWithToken 创建带 AuthHandler 路由的测试引擎,同时返回 TokenManager 以便生成测试 token。
|
||||
func newTestRouterWithToken(svc auth.Service) (*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.NewAuthHandler(svc, tm)
|
||||
h.RegisterRoutes(r.Group("/api"))
|
||||
return r, tm
|
||||
}
|
||||
|
||||
func sampleAuthResponse() *auth.AuthResponse {
|
||||
return &auth.AuthResponse{
|
||||
User: auth.UserResponse{
|
||||
ID: "user-123",
|
||||
Username: "alice",
|
||||
},
|
||||
AccessToken: "access-token",
|
||||
RefreshToken: "refresh-token",
|
||||
}
|
||||
}
|
||||
|
||||
// --- Register ---
|
||||
|
||||
func TestRegister_Success(t *testing.T) {
|
||||
svc := &mockAuthService{
|
||||
RegisterFunc: func(_ context.Context, req auth.RegisterRequest) (*auth.AuthResponse, error) {
|
||||
assert.Equal(t, "alice", req.Username)
|
||||
assert.Equal(t, "password123", req.Password)
|
||||
return sampleAuthResponse(), nil
|
||||
},
|
||||
}
|
||||
r := newTestRouter(svc)
|
||||
|
||||
body, _ := json.Marshal(auth.RegisterRequest{Username: "alice", Password: "password123"})
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/auth/register", bytes.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusCreated, w.Code)
|
||||
var resp auth.AuthResponse
|
||||
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
|
||||
assert.Equal(t, "alice", resp.User.Username)
|
||||
assert.NotEmpty(t, resp.AccessToken)
|
||||
}
|
||||
|
||||
func TestRegister_InvalidInput_EmptyBody(t *testing.T) {
|
||||
svc := &mockAuthService{}
|
||||
r := newTestRouter(svc)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/auth/register", nil)
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
assert.Contains(t, w.Body.String(), "INVALID_INPUT")
|
||||
}
|
||||
|
||||
func TestRegister_InvalidInput_UsernameTooShort(t *testing.T) {
|
||||
svc := &mockAuthService{}
|
||||
r := newTestRouter(svc)
|
||||
|
||||
body, _ := json.Marshal(auth.RegisterRequest{Username: "ab", Password: "password123"})
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/auth/register", bytes.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
assert.Contains(t, w.Body.String(), "username must be 3-64 characters")
|
||||
}
|
||||
|
||||
func TestRegister_InvalidInput_PasswordTooShort(t *testing.T) {
|
||||
svc := &mockAuthService{}
|
||||
r := newTestRouter(svc)
|
||||
|
||||
body, _ := json.Marshal(auth.RegisterRequest{Username: "alice", Password: "short"})
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/auth/register", bytes.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
assert.Contains(t, w.Body.String(), "password must be 8-72 characters")
|
||||
}
|
||||
|
||||
func TestRegister_UsernameTaken(t *testing.T) {
|
||||
svc := &mockAuthService{
|
||||
RegisterFunc: func(_ context.Context, _ auth.RegisterRequest) (*auth.AuthResponse, error) {
|
||||
return nil, auth.ErrUsernameTaken
|
||||
},
|
||||
}
|
||||
r := newTestRouter(svc)
|
||||
|
||||
body, _ := json.Marshal(auth.RegisterRequest{Username: "alice", Password: "password123"})
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/auth/register", bytes.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusConflict, w.Code)
|
||||
assert.Contains(t, w.Body.String(), "USERNAME_TAKEN")
|
||||
}
|
||||
|
||||
// --- Login ---
|
||||
|
||||
func TestLogin_Success(t *testing.T) {
|
||||
svc := &mockAuthService{
|
||||
LoginFunc: func(_ context.Context, req auth.LoginRequest) (*auth.AuthResponse, error) {
|
||||
assert.Equal(t, "alice", req.Username)
|
||||
assert.Equal(t, "password123", req.Password)
|
||||
return sampleAuthResponse(), nil
|
||||
},
|
||||
}
|
||||
r := newTestRouter(svc)
|
||||
|
||||
body, _ := json.Marshal(auth.LoginRequest{Username: "alice", Password: "password123"})
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/auth/login", bytes.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
var resp auth.AuthResponse
|
||||
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
|
||||
assert.Equal(t, "alice", resp.User.Username)
|
||||
}
|
||||
|
||||
func TestLogin_InvalidCredentials(t *testing.T) {
|
||||
svc := &mockAuthService{
|
||||
LoginFunc: func(_ context.Context, _ auth.LoginRequest) (*auth.AuthResponse, error) {
|
||||
return nil, auth.ErrInvalidCredentials
|
||||
},
|
||||
}
|
||||
r := newTestRouter(svc)
|
||||
|
||||
body, _ := json.Marshal(auth.LoginRequest{Username: "alice", Password: "wrong-password"})
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/auth/login", bytes.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusUnauthorized, w.Code)
|
||||
assert.Contains(t, w.Body.String(), "INVALID_CREDENTIALS")
|
||||
}
|
||||
|
||||
// --- Refresh ---
|
||||
|
||||
func TestRefresh_Success(t *testing.T) {
|
||||
svc := &mockAuthService{
|
||||
RefreshFunc: func(_ context.Context, req auth.RefreshRequest) (*auth.AuthResponse, error) {
|
||||
assert.Equal(t, "some-refresh-token", req.RefreshToken)
|
||||
return sampleAuthResponse(), nil
|
||||
},
|
||||
}
|
||||
r := newTestRouter(svc)
|
||||
|
||||
body, _ := json.Marshal(auth.RefreshRequest{RefreshToken: "some-refresh-token"})
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/auth/refresh", bytes.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusOK, w.Code)
|
||||
}
|
||||
|
||||
func TestRefresh_MissingToken(t *testing.T) {
|
||||
svc := &mockAuthService{}
|
||||
r := newTestRouter(svc)
|
||||
|
||||
body, _ := json.Marshal(auth.RefreshRequest{RefreshToken: ""})
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/auth/refresh", bytes.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
assert.Contains(t, w.Body.String(), "refresh_token is required")
|
||||
}
|
||||
|
||||
func TestRefresh_UsedToken(t *testing.T) {
|
||||
svc := &mockAuthService{
|
||||
RefreshFunc: func(_ context.Context, _ auth.RefreshRequest) (*auth.AuthResponse, error) {
|
||||
return nil, auth.ErrRefreshTokenUsed
|
||||
},
|
||||
}
|
||||
r := newTestRouter(svc)
|
||||
|
||||
body, _ := json.Marshal(auth.RefreshRequest{RefreshToken: "used-token"})
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/auth/refresh", bytes.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusUnauthorized, w.Code)
|
||||
assert.Contains(t, w.Body.String(), "INVALID_TOKEN")
|
||||
}
|
||||
|
||||
// --- Logout ---
|
||||
|
||||
func TestLogout_Success(t *testing.T) {
|
||||
logoutCalled := false
|
||||
svc := &mockAuthService{
|
||||
LogoutFunc: func(_ context.Context, userID, refreshToken string) error {
|
||||
assert.Equal(t, "user-123", userID)
|
||||
assert.Equal(t, "refresh-token-to-revoke", refreshToken)
|
||||
logoutCalled = true
|
||||
return nil
|
||||
},
|
||||
}
|
||||
r, tm := newTestRouterWithToken(svc)
|
||||
|
||||
// 生成有效 token
|
||||
access, _, err := tm.GeneratePair("user-123", "alice")
|
||||
require.NoError(t, err)
|
||||
|
||||
body, _ := json.Marshal(map[string]string{"refresh_token": "refresh-token-to-revoke"})
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/auth/logout", 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.True(t, logoutCalled)
|
||||
assert.Contains(t, w.Body.String(), "logged out successfully")
|
||||
}
|
||||
|
||||
func TestLogout_MissingAuth(t *testing.T) {
|
||||
svc := &mockAuthService{}
|
||||
r := newTestRouter(svc)
|
||||
|
||||
body, _ := json.Marshal(map[string]string{"refresh_token": "some-token"})
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/auth/logout", bytes.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
assert.Equal(t, http.StatusUnauthorized, w.Code)
|
||||
}
|
||||
|
||||
func TestLogout_MissingRefreshToken(t *testing.T) {
|
||||
svc := &mockAuthService{}
|
||||
r, tm := newTestRouterWithToken(svc)
|
||||
|
||||
access, _, err := tm.GeneratePair("user-123", "alice")
|
||||
require.NoError(t, err)
|
||||
|
||||
body, _ := json.Marshal(map[string]string{"refresh_token": ""})
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/auth/logout", 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(), "refresh_token is required")
|
||||
}
|
||||
300
backend/internal/api/conversation.go
Normal file
300
backend/internal/api/conversation.go
Normal file
@@ -0,0 +1,300 @@
|
||||
// Package api 提供 REST API 处理函数。
|
||||
package api
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/hhs/camtalk/internal/auth"
|
||||
apperr "github.com/hhs/camtalk/internal/errors"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/session"
|
||||
)
|
||||
|
||||
// ConversationHandler 提供对话相关的 REST 端点。
|
||||
type ConversationHandler struct {
|
||||
sessionMgr session.Manager
|
||||
tokenMgr *auth.TokenManager
|
||||
}
|
||||
|
||||
// NewConversationHandler 创建 ConversationHandler。
|
||||
func NewConversationHandler(sessionMgr session.Manager, tokenMgr *auth.TokenManager) *ConversationHandler {
|
||||
return &ConversationHandler{
|
||||
sessionMgr: sessionMgr,
|
||||
tokenMgr: tokenMgr,
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterRoutes 注册对话相关路由到给定的路由组。所有端点需要认证。
|
||||
func (h *ConversationHandler) RegisterRoutes(rg *gin.RouterGroup) {
|
||||
conv := rg.Group("/conversations", auth.AuthMiddleware(h.tokenMgr))
|
||||
{
|
||||
conv.GET("", h.List)
|
||||
conv.POST("", h.Create)
|
||||
conv.GET("/:id", h.Get)
|
||||
conv.PATCH("/:id", h.UpdateTitle)
|
||||
conv.DELETE("/:id", h.Delete)
|
||||
conv.GET("/:id/messages", h.GetMessages)
|
||||
}
|
||||
}
|
||||
|
||||
// List GET /api/conversations — 获取当前用户的对话列表。
|
||||
func (h *ConversationHandler) List(c *gin.Context) {
|
||||
userID := c.GetString(auth.ContextKeyUserID)
|
||||
|
||||
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
|
||||
size, _ := strconv.Atoi(c.DefaultQuery("size", "20"))
|
||||
|
||||
if page <= 0 {
|
||||
page = 1
|
||||
}
|
||||
if size <= 0 || size > 100 {
|
||||
size = 20
|
||||
}
|
||||
|
||||
summaries, total, err := h.sessionMgr.ListByUser(c.Request.Context(), userID, page, size)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": apperr.CodeInternalError,
|
||||
"message": "failed to list conversations",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"conversations": summaries,
|
||||
"total": total,
|
||||
"page": page,
|
||||
"size": size,
|
||||
})
|
||||
}
|
||||
|
||||
// CreateConversationRequest POST /api/conversations 请求体。
|
||||
type CreateConversationRequest struct {
|
||||
Config *models.SessionConfig `json:"config,omitempty"`
|
||||
}
|
||||
|
||||
// Create POST /api/conversations — 创建新对话。
|
||||
func (h *ConversationHandler) Create(c *gin.Context) {
|
||||
userID := c.GetString(auth.ContextKeyUserID)
|
||||
|
||||
var req CreateConversationRequest
|
||||
_ = c.ShouldBindJSON(&req)
|
||||
|
||||
cfg := models.DefaultConfig()
|
||||
if req.Config != nil {
|
||||
cfg = *req.Config
|
||||
}
|
||||
|
||||
sessionID, err := h.sessionMgr.Create(c.Request.Context(), userID, cfg)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": apperr.CodeInternalError,
|
||||
"message": "failed to create conversation",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
sess, err := h.sessionMgr.Get(c.Request.Context(), sessionID)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": apperr.CodeInternalError,
|
||||
"message": "failed to retrieve created conversation",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusCreated, gin.H{
|
||||
"id": sess.ID,
|
||||
"title": sess.Title,
|
||||
"created_at": sess.CreatedAt,
|
||||
"updated_at": sess.UpdatedAt,
|
||||
})
|
||||
}
|
||||
|
||||
// Get GET /api/conversations/:id — 获取对话详情。
|
||||
func (h *ConversationHandler) Get(c *gin.Context) {
|
||||
sessionID := c.Param("id")
|
||||
|
||||
sess, err := h.getSessionForUser(c, sessionID)
|
||||
if err != nil {
|
||||
return // getSessionForUser 已写入响应
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"id": sess.ID,
|
||||
"title": sess.Title,
|
||||
"created_at": sess.CreatedAt,
|
||||
"updated_at": sess.UpdatedAt,
|
||||
"config": sess.Config,
|
||||
})
|
||||
}
|
||||
|
||||
// UpdateTitleRequest PATCH /api/conversations/:id 请求体。
|
||||
type UpdateTitleRequest struct {
|
||||
Title string `json:"title"`
|
||||
}
|
||||
|
||||
// UpdateTitle PATCH /api/conversations/:id — 更新对话标题。
|
||||
func (h *ConversationHandler) UpdateTitle(c *gin.Context) {
|
||||
sessionID := c.Param("id")
|
||||
|
||||
// 先校验归属
|
||||
if _, err := h.getSessionForUser(c, sessionID); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
var req UpdateTitleRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil || req.Title == "" {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
"code": apperr.CodeInvalidInput,
|
||||
"message": "title is required",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if len([]rune(req.Title)) > 100 {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
"code": apperr.CodeInvalidInput,
|
||||
"message": "title must be 100 characters or less",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.sessionMgr.UpdateTitle(c.Request.Context(), sessionID, req.Title); err != nil {
|
||||
if errors.Is(err, session.ErrSessionNotFound) {
|
||||
c.JSON(http.StatusNotFound, gin.H{
|
||||
"code": apperr.CodeSessionNotFound,
|
||||
"message": "conversation not found",
|
||||
})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": apperr.CodeInternalError,
|
||||
"message": "failed to update title",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"message": "title updated",
|
||||
})
|
||||
}
|
||||
|
||||
// Delete DELETE /api/conversations/:id — 删除对话。
|
||||
func (h *ConversationHandler) Delete(c *gin.Context) {
|
||||
sessionID := c.Param("id")
|
||||
|
||||
// 先校验归属
|
||||
if _, err := h.getSessionForUser(c, sessionID); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.sessionMgr.Destroy(c.Request.Context(), sessionID); err != nil {
|
||||
if errors.Is(err, session.ErrSessionNotFound) {
|
||||
c.JSON(http.StatusNotFound, gin.H{
|
||||
"code": apperr.CodeSessionNotFound,
|
||||
"message": "conversation not found",
|
||||
})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": apperr.CodeInternalError,
|
||||
"message": "failed to delete conversation",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
|
||||
// GetMessages GET /api/conversations/:id/messages — 获取对话消息列表。
|
||||
//
|
||||
// 查询参数:
|
||||
// - limit: 返回消息数量上限,默认 50
|
||||
// - before: 消息偏移量(用于分页),返回此偏移量之前的消息
|
||||
func (h *ConversationHandler) GetMessages(c *gin.Context) {
|
||||
sessionID := c.Param("id")
|
||||
|
||||
// 先校验归属
|
||||
if _, err := h.getSessionForUser(c, sessionID); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
limit, _ := strconv.Atoi(c.DefaultQuery("limit", "50"))
|
||||
if limit <= 0 || limit > 200 {
|
||||
limit = 50
|
||||
}
|
||||
|
||||
before, _ := strconv.Atoi(c.DefaultQuery("before", "0"))
|
||||
|
||||
// 获取全量历史(内存实现中 history 是全量存储的)
|
||||
allMessages, err := h.sessionMgr.GetHistory(c.Request.Context(), sessionID, 0)
|
||||
if err != nil {
|
||||
if errors.Is(err, session.ErrSessionNotFound) {
|
||||
c.JSON(http.StatusNotFound, gin.H{
|
||||
"code": apperr.CodeSessionNotFound,
|
||||
"message": "conversation not found",
|
||||
})
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": apperr.CodeInternalError,
|
||||
"message": "failed to get messages",
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
total := len(allMessages)
|
||||
|
||||
// before > 0 表示取 before 之前的消息(不含 before 位置)
|
||||
if before > 0 && before <= total {
|
||||
allMessages = allMessages[:before]
|
||||
}
|
||||
|
||||
// 取最后 limit 条
|
||||
start := len(allMessages) - limit
|
||||
if start < 0 {
|
||||
start = 0
|
||||
}
|
||||
messages := allMessages[start:]
|
||||
|
||||
c.JSON(http.StatusOK, gin.H{
|
||||
"messages": messages,
|
||||
"total": total,
|
||||
})
|
||||
}
|
||||
|
||||
// getSessionForUser 获取会话并校验当前用户是否有权限访问。
|
||||
// 返回 404(而非 403)以避免信息泄露。
|
||||
func (h *ConversationHandler) getSessionForUser(c *gin.Context, sessionID string) (*models.Session, error) {
|
||||
sess, err := h.sessionMgr.Get(c.Request.Context(), sessionID)
|
||||
if err != nil {
|
||||
if errors.Is(err, session.ErrSessionNotFound) {
|
||||
c.JSON(http.StatusNotFound, gin.H{
|
||||
"code": apperr.CodeSessionNotFound,
|
||||
"message": "conversation not found",
|
||||
})
|
||||
} else {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": apperr.CodeInternalError,
|
||||
"message": "internal server error",
|
||||
})
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
userID := c.GetString(auth.ContextKeyUserID)
|
||||
if sess.UserID != userID {
|
||||
c.JSON(http.StatusNotFound, gin.H{
|
||||
"code": apperr.CodeSessionNotFound,
|
||||
"message": "conversation not found",
|
||||
})
|
||||
return nil, errors.New("forbidden")
|
||||
}
|
||||
|
||||
return sess, nil
|
||||
}
|
||||
575
backend/internal/api/conversation_test.go
Normal file
575
backend/internal/api/conversation_test.go
Normal file
@@ -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)
|
||||
}
|
||||
@@ -36,7 +36,7 @@ func (h *SessionHandler) CreateSession(c *gin.Context) {
|
||||
cfg = *req.Config
|
||||
}
|
||||
|
||||
sessionID, err := h.sessionMgr.Create(c.Request.Context(), cfg)
|
||||
sessionID, err := h.sessionMgr.Create(c.Request.Context(), "", cfg)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{
|
||||
"code": "INTERNAL_ERROR",
|
||||
|
||||
@@ -14,6 +14,12 @@ const (
|
||||
CodeSTTError = "STT_ERROR"
|
||||
CodeTTSError = "TTS_ERROR"
|
||||
CodeInternalError = "INTERNAL_ERROR"
|
||||
|
||||
// 认证相关错误码
|
||||
CodeUsernameTaken = "USERNAME_TAKEN"
|
||||
CodeInvalidCredentials = "INVALID_CREDENTIALS"
|
||||
CodeInvalidToken = "INVALID_TOKEN"
|
||||
CodeInvalidInput = "INVALID_INPUT"
|
||||
)
|
||||
|
||||
// Sender 定义发送 WS 错误消息的接口,便于测试 mock。
|
||||
|
||||
@@ -5,7 +5,10 @@ import "time"
|
||||
// Session 会话。
|
||||
type Session struct {
|
||||
ID string `json:"session_id"`
|
||||
UserID string `json:"user_id,omitempty"`
|
||||
Title string `json:"title"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
Config SessionConfig `json:"config"`
|
||||
}
|
||||
|
||||
@@ -16,6 +19,9 @@ type SessionConfig struct {
|
||||
Language string `json:"language"`
|
||||
}
|
||||
|
||||
// DefaultSessionTitle 默认会话标题。
|
||||
const DefaultSessionTitle = "新对话"
|
||||
|
||||
// DefaultConfig 默认会话配置。
|
||||
func DefaultConfig() SessionConfig {
|
||||
return SessionConfig{TTSEnabled: true, DetailLevel: "low", Language: "zh-CN"}
|
||||
|
||||
@@ -16,6 +16,7 @@ import (
|
||||
"github.com/hhs/camtalk/internal/config"
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/session"
|
||||
)
|
||||
|
||||
func init() {
|
||||
@@ -63,11 +64,21 @@ type MockSessionManager struct {
|
||||
mock.Mock
|
||||
}
|
||||
|
||||
func (m *MockSessionManager) Create(ctx context.Context, config models.SessionConfig) (string, error) {
|
||||
args := m.Called(ctx, config)
|
||||
func (m *MockSessionManager) Create(ctx context.Context, userID string, config models.SessionConfig) (string, error) {
|
||||
args := m.Called(ctx, userID, config)
|
||||
return args.String(0), args.Error(1)
|
||||
}
|
||||
|
||||
func (m *MockSessionManager) UpdateTitle(ctx context.Context, sessionID string, title string) error {
|
||||
args := m.Called(ctx, sessionID, title)
|
||||
return args.Error(0)
|
||||
}
|
||||
|
||||
func (m *MockSessionManager) ListByUser(ctx context.Context, userID string, page, size int) ([]session.ConversationSummary, int, error) {
|
||||
args := m.Called(ctx, userID, page, size)
|
||||
return args.Get(0).([]session.ConversationSummary), args.Int(1), args.Error(2)
|
||||
}
|
||||
|
||||
func (m *MockSessionManager) Get(ctx context.Context, sessionID string) (*models.Session, error) {
|
||||
args := m.Called(ctx, sessionID)
|
||||
if args.Get(0) == nil {
|
||||
|
||||
@@ -4,6 +4,7 @@ package session
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
)
|
||||
@@ -11,11 +12,20 @@ import (
|
||||
// ErrSessionNotFound 会话不存在或已过期。
|
||||
var ErrSessionNotFound = errors.New("session not found")
|
||||
|
||||
// ConversationSummary 对话摘要(列表展示用)。
|
||||
type ConversationSummary struct {
|
||||
ID string `json:"id"`
|
||||
Title string `json:"title"`
|
||||
LastMessage string `json:"last_message"`
|
||||
MessageCount int `json:"message_count"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
// Manager 会话管理器接口。
|
||||
// WebSocket Handler 通过此接口操作会话,不直接接触存储层。
|
||||
type Manager interface {
|
||||
// Create 创建新会话,返回 session ID。
|
||||
Create(ctx context.Context, config models.SessionConfig) (string, error)
|
||||
// Create 创建新会话,返回 session ID。userID 为空表示匿名会话。
|
||||
Create(ctx context.Context, userID string, config models.SessionConfig) (string, error)
|
||||
|
||||
// Get 获取会话(含 config)。不存在返回 ErrSessionNotFound。
|
||||
Get(ctx context.Context, sessionID string) (*models.Session, error)
|
||||
@@ -23,6 +33,12 @@ type Manager interface {
|
||||
// UpdateConfig 更新会话配置(config 消息触发)。
|
||||
UpdateConfig(ctx context.Context, sessionID string, patch models.SessionConfigPatch) error
|
||||
|
||||
// UpdateTitle 更新会话标题。
|
||||
UpdateTitle(ctx context.Context, sessionID string, title string) error
|
||||
|
||||
// ListByUser 获取用户的对话列表(分页,按 UpdatedAt 降序)。
|
||||
ListByUser(ctx context.Context, userID string, page, size int) ([]ConversationSummary, int, error)
|
||||
|
||||
// GetHistory 获取最近 N 轮对话历史(供 Orchestrator 构建 LLM 上下文)。
|
||||
GetHistory(ctx context.Context, sessionID string, limit int) ([]models.Message, error)
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@ package session
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sort"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
@@ -9,6 +10,7 @@ import (
|
||||
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
"github.com/hhs/camtalk/internal/store"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -32,11 +34,23 @@ type MemoryManager struct {
|
||||
ttl time.Duration
|
||||
maxHistory int
|
||||
stopCleaner chan struct{}
|
||||
msgRepo store.MessageRepository // 可选,消息持久化(Write-Through)
|
||||
}
|
||||
|
||||
// Option MemoryManager 的函数式选项。
|
||||
type Option func(*MemoryManager)
|
||||
|
||||
// WithMessageRepository 注入消息持久化仓库,启用 Write-Through 模式。
|
||||
func WithMessageRepository(repo store.MessageRepository) Option {
|
||||
return func(m *MemoryManager) {
|
||||
m.msgRepo = repo
|
||||
}
|
||||
}
|
||||
|
||||
// NewMemoryManager 创建内存版 SessionManager。
|
||||
// ttl 为会话过期时间,maxHistory 为对话历史上限(0 表示使用默认值 20)。
|
||||
func NewMemoryManager(ttl time.Duration, maxHistory int) *MemoryManager {
|
||||
// opts 为可选配置,如 WithMessageRepository 启用消息持久化。
|
||||
func NewMemoryManager(ttl time.Duration, maxHistory int, opts ...Option) *MemoryManager {
|
||||
if ttl <= 0 {
|
||||
ttl = defaultTTL
|
||||
}
|
||||
@@ -51,6 +65,10 @@ func NewMemoryManager(ttl time.Duration, maxHistory int) *MemoryManager {
|
||||
stopCleaner: make(chan struct{}),
|
||||
}
|
||||
|
||||
for _, opt := range opts {
|
||||
opt(m)
|
||||
}
|
||||
|
||||
// 启动后台清理 goroutine,每分钟清除过期会话。
|
||||
go m.cleanLoop()
|
||||
|
||||
@@ -95,8 +113,8 @@ func (m *MemoryManager) isExpired(entry *sessionEntry) bool {
|
||||
return time.Since(entry.lastActive) > m.ttl
|
||||
}
|
||||
|
||||
// Create 创建新会话。
|
||||
func (m *MemoryManager) Create(_ context.Context, config models.SessionConfig) (string, error) {
|
||||
// Create 创建新会话。userID 为空表示匿名会话。
|
||||
func (m *MemoryManager) Create(_ context.Context, userID string, config models.SessionConfig) (string, error) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
@@ -105,14 +123,17 @@ func (m *MemoryManager) Create(_ context.Context, config models.SessionConfig) (
|
||||
m.sessions[id] = &sessionEntry{
|
||||
session: models.Session{
|
||||
ID: id,
|
||||
UserID: userID,
|
||||
Title: models.DefaultSessionTitle,
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
Config: config,
|
||||
},
|
||||
history: make([]models.Message, 0),
|
||||
lastActive: now,
|
||||
}
|
||||
|
||||
logger.Log.Debugw("session created", "session", id)
|
||||
logger.Log.Debugw("session created", "session", id, "user_id", userID)
|
||||
return id, nil
|
||||
}
|
||||
|
||||
@@ -147,6 +168,94 @@ func (m *MemoryManager) UpdateConfig(_ context.Context, sessionID string, patch
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdateTitle 更新会话标题。
|
||||
func (m *MemoryManager) UpdateTitle(_ context.Context, sessionID string, title string) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
entry, ok := m.sessions[sessionID]
|
||||
if !ok || m.isExpired(entry) {
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
|
||||
entry.session.Title = title
|
||||
entry.session.UpdatedAt = time.Now()
|
||||
entry.lastActive = time.Now()
|
||||
|
||||
logger.Log.Debugw("session title updated", "session", sessionID, "title", title)
|
||||
return nil
|
||||
}
|
||||
|
||||
// ListByUser 获取用户的对话列表(分页,按 UpdatedAt 降序)。
|
||||
// 若配置了 MessageRepository,消息统计从 PostgreSQL 聚合查询(更准确)。
|
||||
func (m *MemoryManager) ListByUser(ctx context.Context, userID string, page, size int) ([]ConversationSummary, int, error) {
|
||||
m.mu.RLock()
|
||||
|
||||
// 收集该用户的所有 session
|
||||
var list []ConversationSummary
|
||||
var sessionIDs []string
|
||||
for _, entry := range m.sessions {
|
||||
if entry.session.UserID != userID {
|
||||
continue
|
||||
}
|
||||
if m.isExpired(entry) {
|
||||
continue
|
||||
}
|
||||
summary := ConversationSummary{
|
||||
ID: entry.session.ID,
|
||||
Title: entry.session.Title,
|
||||
UpdatedAt: entry.lastActive,
|
||||
}
|
||||
// 先用内存值填充,后续可能被 PG 统计覆盖
|
||||
summary.MessageCount = len(entry.history)
|
||||
if len(entry.history) > 0 {
|
||||
summary.LastMessage = entry.history[len(entry.history)-1].Content
|
||||
}
|
||||
list = append(list, summary)
|
||||
sessionIDs = append(sessionIDs, entry.session.ID)
|
||||
}
|
||||
m.mu.RUnlock()
|
||||
|
||||
// 若配置了 msgRepo,从 PostgreSQL 获取更准确的消息统计
|
||||
if m.msgRepo != nil && len(sessionIDs) > 0 {
|
||||
if stats, err := m.msgRepo.GetSessionMessageStats(ctx, sessionIDs); err == nil {
|
||||
for i := range list {
|
||||
if s, ok := stats[list[i].ID]; ok {
|
||||
list[i].LastMessage = s.LastMessage
|
||||
list[i].MessageCount = s.MessageCount
|
||||
}
|
||||
}
|
||||
} else {
|
||||
logger.Log.Warnw("get session message stats failed, falling back to in-memory", "error", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 按 UpdatedAt 降序排序
|
||||
sort.Slice(list, func(i, j int) bool {
|
||||
return list[i].UpdatedAt.After(list[j].UpdatedAt)
|
||||
})
|
||||
|
||||
total := len(list)
|
||||
|
||||
// 分页
|
||||
if page <= 0 {
|
||||
page = 1
|
||||
}
|
||||
if size <= 0 {
|
||||
size = 20
|
||||
}
|
||||
start := (page - 1) * size
|
||||
if start >= total {
|
||||
return []ConversationSummary{}, total, nil
|
||||
}
|
||||
end := start + size
|
||||
if end > total {
|
||||
end = total
|
||||
}
|
||||
|
||||
return list[start:end], total, nil
|
||||
}
|
||||
|
||||
// GetHistory 获取最近 N 轮对话历史。
|
||||
func (m *MemoryManager) GetHistory(_ context.Context, sessionID string, limit int) ([]models.Message, error) {
|
||||
m.mu.RLock()
|
||||
@@ -168,26 +277,96 @@ func (m *MemoryManager) GetHistory(_ context.Context, sessionID string, limit in
|
||||
}
|
||||
|
||||
// AppendMessage 追加一条对话消息,同时刷新 TTL。
|
||||
// 若配置了 MessageRepository,消息会异步写入 PostgreSQL(Write-Through)。
|
||||
func (m *MemoryManager) AppendMessage(_ context.Context, sessionID string, msg models.Message) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
entry, ok := m.sessions[sessionID]
|
||||
if !ok || m.isExpired(entry) {
|
||||
m.mu.Unlock()
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
|
||||
entry.history = append(entry.history, msg)
|
||||
|
||||
// 自动更新标题:首条 user 消息时,如果标题为默认值,自动更新为消息前 20 字符
|
||||
if msg.Role == "user" && entry.session.Title == models.DefaultSessionTitle {
|
||||
entry.session.Title = generateTitle(msg.Content)
|
||||
}
|
||||
|
||||
// 超过上限时裁剪,保留最新的 maxHistory 条
|
||||
if len(entry.history) > m.maxHistory {
|
||||
entry.history = entry.history[len(entry.history)-m.maxHistory:]
|
||||
}
|
||||
|
||||
entry.lastActive = time.Now()
|
||||
now := time.Now()
|
||||
entry.lastActive = now
|
||||
entry.session.UpdatedAt = now
|
||||
m.mu.Unlock()
|
||||
|
||||
// Write-Through:异步写冷存储,不阻塞调用方
|
||||
if m.msgRepo != nil {
|
||||
go func() {
|
||||
if err := m.msgRepo.SaveMessage(context.Background(), sessionID, msg, 0); err != nil {
|
||||
logger.Log.Warnw("persist message failed", "session", sessionID, "error", err)
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// generateTitle 从首条消息生成对话标题(取前 20 个字符)。
|
||||
func generateTitle(firstMessage string) string {
|
||||
runes := []rune(firstMessage)
|
||||
if len(runes) > 20 {
|
||||
return string(runes[:20]) + "…"
|
||||
}
|
||||
return firstMessage
|
||||
}
|
||||
|
||||
// LoadSession 从外部存储加载会话到内存热存储。
|
||||
// 用于 conversation_id 恢复场景:WS 连接时会话不在内存中,从 PostgreSQL 加载。
|
||||
// 若会话已在内存中,返回 nil(幂等)。
|
||||
func (m *MemoryManager) LoadSession(sess *models.Session, messages []models.Message) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
if _, ok := m.sessions[sess.ID]; ok {
|
||||
return nil // 已在内存中,无需重复加载
|
||||
}
|
||||
|
||||
m.sessions[sess.ID] = &sessionEntry{
|
||||
session: *sess,
|
||||
history: messages,
|
||||
lastActive: time.Now(),
|
||||
}
|
||||
|
||||
logger.Log.Debugw("session loaded from DB", "session", sess.ID, "messages", len(messages))
|
||||
return nil
|
||||
}
|
||||
|
||||
// LoadSessionFromRepo 从 MessageRepository 加载会话消息并注册到内存。
|
||||
// 适用于已注入 MessageRepository 的场景,调用方只需传入 session 元数据。
|
||||
func (m *MemoryManager) LoadSessionFromRepo(ctx context.Context, sess *models.Session) error {
|
||||
if m.msgRepo == nil {
|
||||
return m.LoadSession(sess, nil)
|
||||
}
|
||||
|
||||
// 从冷存储加载全部消息(limit=0 表示全量)
|
||||
stored, err := m.msgRepo.GetMessages(ctx, sess.ID, 0, 0)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
messages := make([]models.Message, len(stored))
|
||||
for i, s := range stored {
|
||||
messages[i] = models.Message{Role: s.Role, Content: s.Content}
|
||||
}
|
||||
|
||||
return m.LoadSession(sess, messages)
|
||||
}
|
||||
|
||||
// SetActiveRequest 标记当前正在处理的请求 ID。
|
||||
func (m *MemoryManager) SetActiveRequest(_ context.Context, sessionID string, requestID string) error {
|
||||
m.mu.Lock()
|
||||
|
||||
@@ -19,7 +19,7 @@ func TestCreateAndGet(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
config := models.DefaultConfig()
|
||||
id, err := m.Create(ctx, config)
|
||||
id, err := m.Create(ctx, "", config)
|
||||
if err != nil {
|
||||
t.Fatalf("Create: %v", err)
|
||||
}
|
||||
@@ -56,7 +56,7 @@ func TestExpire(t *testing.T) {
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, models.DefaultConfig())
|
||||
id, _ := m.Create(ctx, "", models.DefaultConfig())
|
||||
|
||||
// 未过期时应能获取
|
||||
_, err := m.Get(ctx, id)
|
||||
@@ -78,7 +78,7 @@ func TestDestroy(t *testing.T) {
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, models.DefaultConfig())
|
||||
id, _ := m.Create(ctx, "", models.DefaultConfig())
|
||||
|
||||
if err := m.Destroy(ctx, id); err != nil {
|
||||
t.Fatalf("Destroy: %v", err)
|
||||
@@ -106,7 +106,7 @@ func TestAppendMessageAndGetHistory(t *testing.T) {
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, models.DefaultConfig())
|
||||
id, _ := m.Create(ctx, "", models.DefaultConfig())
|
||||
|
||||
msgs := []models.Message{
|
||||
{Role: "user", Content: "你好"},
|
||||
@@ -138,7 +138,7 @@ func TestGetHistoryLimit(t *testing.T) {
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, models.DefaultConfig())
|
||||
id, _ := m.Create(ctx, "", models.DefaultConfig())
|
||||
|
||||
for i := 0; i < 10; i++ {
|
||||
m.AppendMessage(ctx, id, models.Message{Role: "user", Content: "msg"})
|
||||
@@ -159,7 +159,7 @@ func TestHistoryLimit(t *testing.T) {
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, models.DefaultConfig())
|
||||
id, _ := m.Create(ctx, "", models.DefaultConfig())
|
||||
|
||||
// 插入超过上限的消息
|
||||
for i := 0; i < 10; i++ {
|
||||
@@ -180,7 +180,7 @@ func TestUpdateConfig(t *testing.T) {
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, models.DefaultConfig())
|
||||
id, _ := m.Create(ctx, "", models.DefaultConfig())
|
||||
|
||||
ttsEnabled := false
|
||||
detailLevel := "high"
|
||||
@@ -211,7 +211,7 @@ func TestActiveRequest(t *testing.T) {
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, models.DefaultConfig())
|
||||
id, _ := m.Create(ctx, "", models.DefaultConfig())
|
||||
|
||||
// 初始应为空
|
||||
reqID, err := m.GetActiveRequestID(ctx, id)
|
||||
@@ -246,7 +246,7 @@ func TestTouchRefreshesTTL(t *testing.T) {
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, models.DefaultConfig())
|
||||
id, _ := m.Create(ctx, "", models.DefaultConfig())
|
||||
|
||||
// 50ms 后 Touch,应重置 TTL
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
@@ -278,9 +278,204 @@ func TestActiveCount(t *testing.T) {
|
||||
t.Errorf("initial ActiveCount = %d, want 0", m.ActiveCount())
|
||||
}
|
||||
|
||||
m.Create(ctx, models.DefaultConfig())
|
||||
m.Create(ctx, models.DefaultConfig())
|
||||
m.Create(ctx, "", models.DefaultConfig())
|
||||
m.Create(ctx, "", models.DefaultConfig())
|
||||
if m.ActiveCount() != 2 {
|
||||
t.Errorf("ActiveCount = %d, want 2", m.ActiveCount())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateWithUserID(t *testing.T) {
|
||||
m := NewMemoryManager(30*time.Minute, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, err := m.Create(ctx, "user-123", models.DefaultConfig())
|
||||
if err != nil {
|
||||
t.Fatalf("Create: %v", err)
|
||||
}
|
||||
|
||||
sess, err := m.Get(ctx, id)
|
||||
if err != nil {
|
||||
t.Fatalf("Get: %v", err)
|
||||
}
|
||||
if sess.UserID != "user-123" {
|
||||
t.Errorf("UserID = %q, want %q", sess.UserID, "user-123")
|
||||
}
|
||||
if sess.Title != models.DefaultSessionTitle {
|
||||
t.Errorf("Title = %q, want %q", sess.Title, models.DefaultSessionTitle)
|
||||
}
|
||||
if sess.UpdatedAt.IsZero() {
|
||||
t.Error("UpdatedAt should not be zero")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateTitle(t *testing.T) {
|
||||
m := NewMemoryManager(30*time.Minute, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, "user-1", models.DefaultConfig())
|
||||
|
||||
if err := m.UpdateTitle(ctx, id, "自定义标题"); err != nil {
|
||||
t.Fatalf("UpdateTitle: %v", err)
|
||||
}
|
||||
|
||||
sess, _ := m.Get(ctx, id)
|
||||
if sess.Title != "自定义标题" {
|
||||
t.Errorf("Title = %q, want %q", sess.Title, "自定义标题")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateTitleNotFound(t *testing.T) {
|
||||
m := NewMemoryManager(30*time.Minute, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
err := m.UpdateTitle(ctx, "nonexistent", "标题")
|
||||
if err != ErrSessionNotFound {
|
||||
t.Errorf("UpdateTitle nonexistent: err = %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAutoTitleOnFirstMessage(t *testing.T) {
|
||||
m := NewMemoryManager(30*time.Minute, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, "user-1", models.DefaultConfig())
|
||||
|
||||
// 首条 user 消息应自动更新标题
|
||||
m.AppendMessage(ctx, id, models.Message{Role: "user", Content: "你好世界"})
|
||||
|
||||
sess, _ := m.Get(ctx, id)
|
||||
if sess.Title != "你好世界" {
|
||||
t.Errorf("Title = %q, want %q", sess.Title, "你好世界")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAutoTitleLongMessage(t *testing.T) {
|
||||
m := NewMemoryManager(30*time.Minute, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, "user-1", models.DefaultConfig())
|
||||
|
||||
// 超过 20 字符的消息应截断
|
||||
longMsg := "这是一条很长很长很长很长很长很长很长很长的消息"
|
||||
m.AppendMessage(ctx, id, models.Message{Role: "user", Content: longMsg})
|
||||
|
||||
sess, _ := m.Get(ctx, id)
|
||||
expected := string([]rune(longMsg)[:20]) + "…"
|
||||
if sess.Title != expected {
|
||||
t.Errorf("Title = %q, want %q", sess.Title, expected)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAutoTitleNotOverwritten(t *testing.T) {
|
||||
m := NewMemoryManager(30*time.Minute, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
id, _ := m.Create(ctx, "user-1", models.DefaultConfig())
|
||||
|
||||
// 首条消息设置标题
|
||||
m.AppendMessage(ctx, id, models.Message{Role: "user", Content: "第一条消息"})
|
||||
// 第二条消息不应覆盖已有的标题
|
||||
m.AppendMessage(ctx, id, models.Message{Role: "user", Content: "第二条消息"})
|
||||
|
||||
sess, _ := m.Get(ctx, id)
|
||||
if sess.Title != "第一条消息" {
|
||||
t.Errorf("Title = %q, want %q", sess.Title, "第一条消息")
|
||||
}
|
||||
}
|
||||
|
||||
func TestListByUser(t *testing.T) {
|
||||
m := NewMemoryManager(30*time.Minute, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
// 创建两个用户的不同会话
|
||||
id1, _ := m.Create(ctx, "user-1", models.DefaultConfig())
|
||||
m.AppendMessage(ctx, id1, models.Message{Role: "user", Content: "会话1"})
|
||||
id2, _ := m.Create(ctx, "user-1", models.DefaultConfig())
|
||||
m.AppendMessage(ctx, id2, models.Message{Role: "user", Content: "会话2"})
|
||||
m.Create(ctx, "user-2", models.DefaultConfig()) // 其他用户的会话
|
||||
|
||||
list, total, err := m.ListByUser(ctx, "user-1", 1, 10)
|
||||
if err != nil {
|
||||
t.Fatalf("ListByUser: %v", err)
|
||||
}
|
||||
if total != 2 {
|
||||
t.Errorf("total = %d, want 2", total)
|
||||
}
|
||||
if len(list) != 2 {
|
||||
t.Fatalf("len = %d, want 2", len(list))
|
||||
}
|
||||
// 按 UpdatedAt 降序,id2 应在前
|
||||
if list[0].ID != id2 {
|
||||
t.Errorf("list[0].ID = %q, want %q", list[0].ID, id2)
|
||||
}
|
||||
if list[0].Title != "会话2" {
|
||||
t.Errorf("list[0].Title = %q, want %q", list[0].Title, "会话2")
|
||||
}
|
||||
if list[0].LastMessage != "会话2" {
|
||||
t.Errorf("list[0].LastMessage = %q, want %q", list[0].LastMessage, "会话2")
|
||||
}
|
||||
}
|
||||
|
||||
func TestListByUserPagination(t *testing.T) {
|
||||
m := NewMemoryManager(30*time.Minute, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
// 创建 5 个会话
|
||||
for i := 0; i < 5; i++ {
|
||||
id, _ := m.Create(ctx, "user-1", models.DefaultConfig())
|
||||
m.AppendMessage(ctx, id, models.Message{Role: "user", Content: "msg"})
|
||||
}
|
||||
|
||||
// 第 1 页,每页 2 条
|
||||
list, total, _ := m.ListByUser(ctx, "user-1", 1, 2)
|
||||
if total != 5 {
|
||||
t.Errorf("total = %d, want 5", total)
|
||||
}
|
||||
if len(list) != 2 {
|
||||
t.Errorf("page 1 len = %d, want 2", len(list))
|
||||
}
|
||||
|
||||
// 第 2 页
|
||||
list, _, _ = m.ListByUser(ctx, "user-1", 2, 2)
|
||||
if len(list) != 2 {
|
||||
t.Errorf("page 2 len = %d, want 2", len(list))
|
||||
}
|
||||
|
||||
// 第 3 页(最后一页)
|
||||
list, _, _ = m.ListByUser(ctx, "user-1", 3, 2)
|
||||
if len(list) != 1 {
|
||||
t.Errorf("page 3 len = %d, want 1", len(list))
|
||||
}
|
||||
|
||||
// 超出范围的页
|
||||
list, _, _ = m.ListByUser(ctx, "user-1", 10, 2)
|
||||
if len(list) != 0 {
|
||||
t.Errorf("out of range page len = %d, want 0", len(list))
|
||||
}
|
||||
}
|
||||
|
||||
func TestListByUserEmpty(t *testing.T) {
|
||||
m := NewMemoryManager(30*time.Minute, 20)
|
||||
defer m.Stop()
|
||||
ctx := context.Background()
|
||||
|
||||
list, total, err := m.ListByUser(ctx, "no-such-user", 1, 10)
|
||||
if err != nil {
|
||||
t.Fatalf("ListByUser: %v", err)
|
||||
}
|
||||
if total != 0 {
|
||||
t.Errorf("total = %d, want 0", total)
|
||||
}
|
||||
if len(list) != 0 {
|
||||
t.Errorf("len = %d, want 0", len(list))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -18,6 +18,7 @@ import (
|
||||
// 数据结构:
|
||||
// - session:{id}:meta → Hash(会话元数据)
|
||||
// - session:{id}:history → List(对话历史)
|
||||
// - user:{id}:sessions → Set(用户会话索引)
|
||||
type RedisManager struct {
|
||||
rdb *redis.Client
|
||||
ttl time.Duration
|
||||
@@ -35,37 +36,48 @@ func NewRedisManager(rdb *redis.Client, ttl time.Duration, maxHistory int) *Redi
|
||||
return &RedisManager{rdb: rdb, ttl: ttl, maxHistory: maxHistory}
|
||||
}
|
||||
|
||||
func metaKey(id string) string { return fmt.Sprintf("session:%s:meta", id) }
|
||||
func histKey(id string) string { return fmt.Sprintf("session:%s:history", id) }
|
||||
func metaKey(id string) string { return fmt.Sprintf("session:%s:meta", id) }
|
||||
func histKey(id string) string { return fmt.Sprintf("session:%s:history", id) }
|
||||
func userSessKey(id string) string { return fmt.Sprintf("user:%s:sessions", id) }
|
||||
|
||||
// Create 创建新会话。
|
||||
func (m *RedisManager) Create(ctx context.Context, config models.SessionConfig) (string, error) {
|
||||
// Create 创建新会话。userID 为空表示匿名会话。
|
||||
func (m *RedisManager) Create(ctx context.Context, userID string, config models.SessionConfig) (string, error) {
|
||||
id := uuidNew()
|
||||
now := time.Now().UTC()
|
||||
|
||||
pipe := m.rdb.Pipeline()
|
||||
|
||||
// 写入 meta Hash
|
||||
pipe.HSet(ctx, metaKey(id), map[string]interface{}{
|
||||
"session_id": id,
|
||||
"config.tts_enabled": strconv.FormatBool(config.TTSEnabled),
|
||||
meta := map[string]interface{}{
|
||||
"session_id": id,
|
||||
"user_id": userID,
|
||||
"title": models.DefaultSessionTitle,
|
||||
"config.tts_enabled": strconv.FormatBool(config.TTSEnabled),
|
||||
"config.detail_level": config.DetailLevel,
|
||||
"config.language": config.Language,
|
||||
"created_at": now.Format(time.RFC3339),
|
||||
"last_active": now.Format(time.RFC3339),
|
||||
"active_request_id": "",
|
||||
})
|
||||
"config.language": config.Language,
|
||||
"created_at": now.Format(time.RFC3339),
|
||||
"updated_at": now.Format(time.RFC3339),
|
||||
"last_active": now.Format(time.RFC3339),
|
||||
"active_request_id": "",
|
||||
}
|
||||
pipe.HSet(ctx, metaKey(id), meta)
|
||||
pipe.Expire(ctx, metaKey(id), m.ttl)
|
||||
|
||||
// 初始化空 history List
|
||||
pipe.RPush(ctx, histKey(id), placeholderHistoryMark)
|
||||
pipe.Expire(ctx, histKey(id), m.ttl)
|
||||
|
||||
// 如果有 userID,添加到用户会话索引
|
||||
if userID != "" {
|
||||
pipe.SAdd(ctx, userSessKey(userID), id)
|
||||
pipe.Expire(ctx, userSessKey(userID), m.ttl)
|
||||
}
|
||||
|
||||
if _, err := pipe.Exec(ctx); err != nil {
|
||||
return "", fmt.Errorf("redis create session: %w", err)
|
||||
}
|
||||
|
||||
logger.Log.Debugw("redis session created", "session", id)
|
||||
logger.Log.Debugw("redis session created", "session", id, "user_id", userID)
|
||||
return id, nil
|
||||
}
|
||||
|
||||
@@ -83,9 +95,12 @@ func (m *RedisManager) Get(ctx context.Context, sessionID string) (*models.Sessi
|
||||
}
|
||||
|
||||
sess := &models.Session{
|
||||
ID: vals["session_id"],
|
||||
ID: vals["session_id"],
|
||||
UserID: vals["user_id"],
|
||||
Title: vals["title"],
|
||||
}
|
||||
sess.CreatedAt, _ = time.Parse(time.RFC3339, vals["created_at"])
|
||||
sess.UpdatedAt, _ = time.Parse(time.RFC3339, vals["updated_at"])
|
||||
sess.Config.TTSEnabled, _ = strconv.ParseBool(vals["config.tts_enabled"])
|
||||
sess.Config.DetailLevel = vals["config.detail_level"]
|
||||
sess.Config.Language = vals["config.language"]
|
||||
@@ -104,8 +119,10 @@ func (m *RedisManager) UpdateConfig(ctx context.Context, sessionID string, patch
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
|
||||
now := time.Now().UTC().Format(time.RFC3339)
|
||||
fields := map[string]interface{}{
|
||||
"last_active": time.Now().UTC().Format(time.RFC3339),
|
||||
"last_active": now,
|
||||
"updated_at": now,
|
||||
}
|
||||
if patch.TTSEnabled != nil {
|
||||
fields["config.tts_enabled"] = strconv.FormatBool(*patch.TTSEnabled)
|
||||
@@ -127,6 +144,109 @@ func (m *RedisManager) UpdateConfig(ctx context.Context, sessionID string, patch
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdateTitle 更新会话标题。
|
||||
func (m *RedisManager) UpdateTitle(ctx context.Context, sessionID string, title string) error {
|
||||
exists, err := m.rdb.Exists(ctx, metaKey(sessionID)).Result()
|
||||
if err != nil {
|
||||
return fmt.Errorf("redis check session: %w", err)
|
||||
}
|
||||
if exists == 0 {
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
|
||||
now := time.Now().UTC().Format(time.RFC3339)
|
||||
if err := m.rdb.HSet(ctx, metaKey(sessionID), "title", title, "updated_at", now, "last_active", now).Err(); err != nil {
|
||||
return fmt.Errorf("redis update title: %w", err)
|
||||
}
|
||||
|
||||
m.rdb.Expire(ctx, metaKey(sessionID), m.ttl)
|
||||
logger.Log.Debugw("redis session title updated", "session", sessionID, "title", title)
|
||||
return nil
|
||||
}
|
||||
|
||||
// ListByUser 获取用户的对话列表(分页,按 UpdatedAt 降序)。
|
||||
func (m *RedisManager) ListByUser(ctx context.Context, userID string, page, size int) ([]ConversationSummary, int, error) {
|
||||
if page <= 0 {
|
||||
page = 1
|
||||
}
|
||||
if size <= 0 {
|
||||
size = 20
|
||||
}
|
||||
|
||||
// 从用户会话索引获取所有 session ID
|
||||
sessionIDs, err := m.rdb.SMembers(ctx, userSessKey(userID)).Result()
|
||||
if err != nil {
|
||||
return nil, 0, fmt.Errorf("redis list user sessions: %w", err)
|
||||
}
|
||||
|
||||
// 收集有效的会话摘要
|
||||
var list []ConversationSummary
|
||||
for _, sid := range sessionIDs {
|
||||
vals, err := m.rdb.HGetAll(ctx, metaKey(sid)).Result()
|
||||
if err != nil || len(vals) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
updatedAt, _ := time.Parse(time.RFC3339, vals["updated_at"])
|
||||
lastActive, _ := time.Parse(time.RFC3339, vals["last_active"])
|
||||
|
||||
// 检查是否过期
|
||||
if time.Since(lastActive) > m.ttl {
|
||||
continue
|
||||
}
|
||||
|
||||
// 获取最后一条消息
|
||||
lastMsg := ""
|
||||
msgCount := 0
|
||||
raws, err := m.rdb.LRange(ctx, histKey(sid), 0, 0).Result()
|
||||
if err == nil && len(raws) > 0 && raws[0] != placeholderHistoryMark {
|
||||
var msg models.Message
|
||||
if json.Unmarshal([]byte(raws[0]), &msg) == nil {
|
||||
lastMsg = msg.Content
|
||||
}
|
||||
}
|
||||
// 获取消息总数(减去占位符)
|
||||
totalLen, err := m.rdb.LLen(ctx, histKey(sid)).Result()
|
||||
if err == nil {
|
||||
msgCount = int(totalLen)
|
||||
if msgCount > 0 {
|
||||
msgCount-- // 减去占位符
|
||||
}
|
||||
}
|
||||
|
||||
list = append(list, ConversationSummary{
|
||||
ID: vals["session_id"],
|
||||
Title: vals["title"],
|
||||
LastMessage: lastMsg,
|
||||
MessageCount: msgCount,
|
||||
UpdatedAt: updatedAt,
|
||||
})
|
||||
}
|
||||
|
||||
// 按 UpdatedAt 降序排序
|
||||
for i := 0; i < len(list); i++ {
|
||||
for j := i + 1; j < len(list); j++ {
|
||||
if list[j].UpdatedAt.After(list[i].UpdatedAt) {
|
||||
list[i], list[j] = list[j], list[i]
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
total := len(list)
|
||||
|
||||
// 分页
|
||||
start := (page - 1) * size
|
||||
if start >= total {
|
||||
return []ConversationSummary{}, total, nil
|
||||
}
|
||||
end := start + size
|
||||
if end > total {
|
||||
end = total
|
||||
}
|
||||
|
||||
return list[start:end], total, nil
|
||||
}
|
||||
|
||||
// GetHistory 获取最近 N 轮对话历史。
|
||||
func (m *RedisManager) GetHistory(ctx context.Context, sessionID string, limit int) ([]models.Message, error) {
|
||||
// 检查会话是否存在
|
||||
@@ -193,8 +313,18 @@ func (m *RedisManager) AppendMessage(ctx context.Context, sessionID string, msg
|
||||
// 刷新 TTL
|
||||
pipe.Expire(ctx, histKey(sessionID), m.ttl)
|
||||
pipe.Expire(ctx, metaKey(sessionID), m.ttl)
|
||||
// 更新 last_active
|
||||
pipe.HSet(ctx, metaKey(sessionID), "last_active", time.Now().UTC().Format(time.RFC3339))
|
||||
|
||||
now := time.Now().UTC().Format(time.RFC3339)
|
||||
// 更新 last_active 和 updated_at
|
||||
pipe.HSet(ctx, metaKey(sessionID), "last_active", now, "updated_at", now)
|
||||
|
||||
// 自动更新标题:首条 user 消息时,如果标题为默认值
|
||||
if msg.Role == "user" {
|
||||
title, _ := m.rdb.HGet(ctx, metaKey(sessionID), "title").Result()
|
||||
if title == models.DefaultSessionTitle {
|
||||
pipe.HSet(ctx, metaKey(sessionID), "title", generateTitle(msg.Content))
|
||||
}
|
||||
}
|
||||
|
||||
if _, err := pipe.Exec(ctx); err != nil {
|
||||
return fmt.Errorf("redis append message: %w", err)
|
||||
@@ -280,6 +410,9 @@ func (m *RedisManager) Touch(ctx context.Context, sessionID string) error {
|
||||
|
||||
// Destroy 显式销毁会话。
|
||||
func (m *RedisManager) Destroy(ctx context.Context, sessionID string) error {
|
||||
// 先获取 user_id 以便清理索引
|
||||
userID, _ := m.rdb.HGet(ctx, metaKey(sessionID), "user_id").Result()
|
||||
|
||||
deleted, err := m.rdb.Del(ctx, metaKey(sessionID), histKey(sessionID)).Result()
|
||||
if err != nil {
|
||||
return fmt.Errorf("redis destroy session: %w", err)
|
||||
@@ -288,6 +421,11 @@ func (m *RedisManager) Destroy(ctx context.Context, sessionID string) error {
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
|
||||
// 清理用户会话索引
|
||||
if userID != "" {
|
||||
m.rdb.SRem(ctx, userSessKey(userID), sessionID)
|
||||
}
|
||||
|
||||
logger.Log.Debugw("redis session destroyed", "session", sessionID)
|
||||
return nil
|
||||
}
|
||||
|
||||
50
backend/internal/store/message.go
Normal file
50
backend/internal/store/message.go
Normal file
@@ -0,0 +1,50 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
)
|
||||
|
||||
var (
|
||||
// ErrMessageNotFound 消息不存在。
|
||||
ErrMessageNotFound = errors.New("message not found")
|
||||
)
|
||||
|
||||
// MessageRepository 消息持久化接口。
|
||||
type MessageRepository interface {
|
||||
// SaveMessage 保存一条消息。
|
||||
SaveMessage(ctx context.Context, sessionID string, msg models.Message, tokensUsed int) error
|
||||
|
||||
// GetMessages 获取会话的消息列表(分页,按 created_at 升序)。
|
||||
// beforeID 为 0 时从最新开始查询。
|
||||
GetMessages(ctx context.Context, sessionID string, limit int, beforeID int64) ([]StoredMessage, error)
|
||||
|
||||
// GetLastMessage 获取会话的最后一条消息。
|
||||
GetLastMessage(ctx context.Context, sessionID string) (*StoredMessage, error)
|
||||
|
||||
// GetMessageCount 获取会话的消息总数。
|
||||
GetMessageCount(ctx context.Context, sessionID string) (int, error)
|
||||
|
||||
// GetSessionMessageStats 批量查询多个会话的消息统计(last_message + message_count)。
|
||||
// 返回的 map key 为 sessionID,仅包含有消息的会话。
|
||||
GetSessionMessageStats(ctx context.Context, sessionIDs []string) (map[string]SessionMessageStats, error)
|
||||
}
|
||||
|
||||
// SessionMessageStats 单个会话的消息统计(SQL 聚合查询结果)。
|
||||
type SessionMessageStats struct {
|
||||
LastMessage string
|
||||
MessageCount int
|
||||
}
|
||||
|
||||
// StoredMessage 持久化消息模型(store 层)。
|
||||
type StoredMessage struct {
|
||||
ID int64 `json:"id"`
|
||||
SessionID string `json:"-"`
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
TokensUsed int `json:"tokens_used"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
163
backend/internal/store/message_pg.go
Normal file
163
backend/internal/store/message_pg.go
Normal file
@@ -0,0 +1,163 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
)
|
||||
|
||||
// PgMessageRepository 基于 PostgreSQL 的 MessageRepository 实现。
|
||||
type PgMessageRepository struct {
|
||||
pool *pgxpool.Pool
|
||||
}
|
||||
|
||||
// NewPgMessageRepository 创建 PgMessageRepository。
|
||||
func NewPgMessageRepository(pool *pgxpool.Pool) *PgMessageRepository {
|
||||
return &PgMessageRepository{pool: pool}
|
||||
}
|
||||
|
||||
func (r *PgMessageRepository) SaveMessage(ctx context.Context, sessionID string, msg models.Message, tokensUsed int) error {
|
||||
_, err := r.pool.Exec(ctx,
|
||||
`INSERT INTO messages (session_id, role, content, tokens_used) VALUES ($1, $2, $3, $4)`,
|
||||
sessionID, msg.Role, msg.Content, tokensUsed,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *PgMessageRepository) GetMessages(ctx context.Context, sessionID string, limit int, beforeID int64) ([]StoredMessage, error) {
|
||||
if limit <= 0 {
|
||||
limit = 50
|
||||
}
|
||||
|
||||
var rows []StoredMessage
|
||||
var err error
|
||||
|
||||
if beforeID > 0 {
|
||||
rows, err = r.queryMessages(ctx,
|
||||
`SELECT id, session_id, role, content, tokens_used, created_at
|
||||
FROM messages
|
||||
WHERE session_id = $1 AND id < $2
|
||||
ORDER BY id DESC
|
||||
LIMIT $3`,
|
||||
sessionID, beforeID, limit,
|
||||
)
|
||||
} else {
|
||||
rows, err = r.queryMessages(ctx,
|
||||
`SELECT id, session_id, role, content, tokens_used, created_at
|
||||
FROM messages
|
||||
WHERE session_id = $1
|
||||
ORDER BY id DESC
|
||||
LIMIT $2`,
|
||||
sessionID, limit,
|
||||
)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// 反转为升序
|
||||
for i, j := 0, len(rows)-1; i < j; i, j = i+1, j-1 {
|
||||
rows[i], rows[j] = rows[j], rows[i]
|
||||
}
|
||||
|
||||
return rows, nil
|
||||
}
|
||||
|
||||
func (r *PgMessageRepository) queryMessages(ctx context.Context, query string, args ...any) ([]StoredMessage, error) {
|
||||
pgxRows, err := r.pool.Query(ctx, query, args...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer pgxRows.Close()
|
||||
|
||||
var messages []StoredMessage
|
||||
for pgxRows.Next() {
|
||||
var m StoredMessage
|
||||
if err := pgxRows.Scan(&m.ID, &m.SessionID, &m.Role, &m.Content, &m.TokensUsed, &m.CreatedAt); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
messages = append(messages, m)
|
||||
}
|
||||
if err := pgxRows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return messages, nil
|
||||
}
|
||||
|
||||
func (r *PgMessageRepository) GetLastMessage(ctx context.Context, sessionID string) (*StoredMessage, error) {
|
||||
var m StoredMessage
|
||||
err := r.pool.QueryRow(ctx,
|
||||
`SELECT id, session_id, role, content, tokens_used, created_at
|
||||
FROM messages
|
||||
WHERE session_id = $1
|
||||
ORDER BY id DESC
|
||||
LIMIT 1`,
|
||||
sessionID,
|
||||
).Scan(&m.ID, &m.SessionID, &m.Role, &m.Content, &m.TokensUsed, &m.CreatedAt)
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return nil, ErrMessageNotFound
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &m, nil
|
||||
}
|
||||
|
||||
func (r *PgMessageRepository) GetMessageCount(ctx context.Context, sessionID string) (int, error) {
|
||||
var count int
|
||||
err := r.pool.QueryRow(ctx,
|
||||
`SELECT COUNT(*) FROM messages WHERE session_id = $1`,
|
||||
sessionID,
|
||||
).Scan(&count)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
func (r *PgMessageRepository) GetSessionMessageStats(ctx context.Context, sessionIDs []string) (map[string]SessionMessageStats, error) {
|
||||
if len(sessionIDs) == 0 {
|
||||
return map[string]SessionMessageStats{}, nil
|
||||
}
|
||||
|
||||
rows, err := r.pool.Query(ctx,
|
||||
`WITH stats AS (
|
||||
SELECT session_id, COUNT(*) AS cnt
|
||||
FROM messages
|
||||
WHERE session_id = ANY($1)
|
||||
GROUP BY session_id
|
||||
),
|
||||
last_msg AS (
|
||||
SELECT DISTINCT ON (session_id) session_id, content
|
||||
FROM messages
|
||||
WHERE session_id = ANY($1)
|
||||
ORDER BY session_id, id DESC
|
||||
)
|
||||
SELECT s.session_id, s.cnt, COALESCE(lm.content, '')
|
||||
FROM stats s
|
||||
LEFT JOIN last_msg lm ON lm.session_id = s.session_id`,
|
||||
sessionIDs,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
result := make(map[string]SessionMessageStats)
|
||||
for rows.Next() {
|
||||
var sid string
|
||||
var stats SessionMessageStats
|
||||
if err := rows.Scan(&sid, &stats.MessageCount, &stats.LastMessage); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result[sid] = stats
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gorilla/websocket"
|
||||
|
||||
"github.com/hhs/camtalk/internal/auth"
|
||||
"github.com/hhs/camtalk/internal/config"
|
||||
"github.com/hhs/camtalk/internal/errors"
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
@@ -91,7 +92,7 @@ func (w *WSClient) SendError(err models.WsError) error {
|
||||
}
|
||||
|
||||
// ServeWS 处理 WebSocket 升级请求。
|
||||
func ServeWS(sessionMgr session.Manager, orch orchestrator.Orchestrator, cfg *config.Config) gin.HandlerFunc {
|
||||
func ServeWS(sessionMgr session.Manager, orch orchestrator.Orchestrator, cfg *config.Config, tokenMgr *auth.TokenManager) gin.HandlerFunc {
|
||||
upgrader := newUpgrader(cfg)
|
||||
heartbeatInterval := time.Duration(cfg.Server.HeartbeatInterval) * time.Second
|
||||
heartbeatTimeout := time.Duration(cfg.Server.HeartbeatTimeout) * time.Second
|
||||
@@ -100,12 +101,37 @@ func ServeWS(sessionMgr session.Manager, orch orchestrator.Orchestrator, cfg *co
|
||||
maxHistory := cfg.Session.MaxHistory
|
||||
|
||||
return func(c *gin.Context) {
|
||||
serveWS(c, sessionMgr, orch, upgrader, heartbeatInterval, heartbeatTimeout, version, maxHistory)
|
||||
serveWS(c, sessionMgr, orch, upgrader, heartbeatInterval, heartbeatTimeout, version, maxHistory, tokenMgr)
|
||||
}
|
||||
}
|
||||
|
||||
func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orchestrator,
|
||||
upgrader websocket.Upgrader, heartbeatInterval, heartbeatTimeout time.Duration, version string, maxHistory int) {
|
||||
upgrader websocket.Upgrader, heartbeatInterval, heartbeatTimeout time.Duration, version string, maxHistory int, tokenMgr *auth.TokenManager) {
|
||||
|
||||
// --- JWT 认证(upgrade 前完成,失败直接返回 HTTP 错误) ---
|
||||
token := c.Query("token")
|
||||
if token == "" {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"error": "missing token"})
|
||||
return
|
||||
}
|
||||
claims, err := tokenMgr.ValidateAccess(token)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"error": "invalid token"})
|
||||
return
|
||||
}
|
||||
userID := claims.UserID
|
||||
username := claims.Username
|
||||
|
||||
// --- conversation_id 处理(upgrade 前校验归属) ---
|
||||
conversationID := c.Query("conversation_id")
|
||||
if conversationID != "" {
|
||||
sess, err := sessionMgr.Get(c.Request.Context(), conversationID)
|
||||
if err != nil || sess.UserID != userID {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"error": "SESSION_NOT_FOUND"})
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
conn, err := upgrader.Upgrade(c.Writer, c.Request, nil)
|
||||
if err != nil {
|
||||
logger.Log.Errorw("websocket upgrade failed", "error", err)
|
||||
@@ -113,11 +139,17 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
// 创建会话
|
||||
sessionID, err := sessionMgr.Create(context.Background(), models.DefaultConfig())
|
||||
if err != nil {
|
||||
logger.Log.Errorw("create session failed", "error", err)
|
||||
return
|
||||
// 创建或复用会话
|
||||
var sessionID string
|
||||
if conversationID != "" {
|
||||
sessionID = conversationID
|
||||
logger.Log.Infow("resuming conversation", "session", sessionID, "user_id", userID)
|
||||
} else {
|
||||
sessionID, err = sessionMgr.Create(context.Background(), userID, models.DefaultConfig())
|
||||
if err != nil {
|
||||
logger.Log.Errorw("create session failed", "error", err)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
client := &Client{
|
||||
@@ -134,7 +166,7 @@ func serveWS(c *gin.Context, sessionMgr session.Manager, orch orchestrator.Orche
|
||||
SessionID: sessionID,
|
||||
ServerVersion: version,
|
||||
})
|
||||
logger.Log.Infow("client connected", "session", sessionID)
|
||||
logger.Log.Infow("client connected", "session", sessionID, "user_id", userID, "username", username)
|
||||
|
||||
// 心跳检测
|
||||
lastPong := time.Now()
|
||||
|
||||
@@ -2,6 +2,7 @@ package ws
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
@@ -13,6 +14,7 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
"context"
|
||||
|
||||
"github.com/hhs/camtalk/internal/auth"
|
||||
"github.com/hhs/camtalk/internal/config"
|
||||
"github.com/hhs/camtalk/internal/logger"
|
||||
"github.com/hhs/camtalk/internal/models"
|
||||
@@ -132,24 +134,29 @@ func (m *MockOrchestrator) ProcessQuery(
|
||||
// --- 测试辅助函数 ---
|
||||
|
||||
// setupTestServer 创建测试用 Gin 服务器和 WebSocket URL。
|
||||
// 返回的 wsURL 已包含有效 token,可直接连接。
|
||||
func setupTestServer(t *testing.T, orch orchestrator.Orchestrator) (*httptest.Server, string) {
|
||||
t.Helper()
|
||||
|
||||
sessionMgr := session.NewMemoryManager(5*time.Minute, 20)
|
||||
t.Cleanup(func() { sessionMgr.Stop() })
|
||||
|
||||
tokenMgr := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour)
|
||||
|
||||
r := gin.New()
|
||||
cfg := &config.Config{
|
||||
App: config.AppConfig{Version: "test"},
|
||||
Server: config.ServerConfig{HeartbeatInterval: 30, HeartbeatTimeout: 60},
|
||||
Session: config.SessionConfig{MaxHistory: 20},
|
||||
}
|
||||
r.GET("/ws", ServeWS(sessionMgr, orch, cfg))
|
||||
r.GET("/ws", ServeWS(sessionMgr, orch, cfg, tokenMgr))
|
||||
|
||||
srv := httptest.NewServer(r)
|
||||
|
||||
// 构造 WebSocket URL
|
||||
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws"
|
||||
// 生成有效 token 并构造 WebSocket URL
|
||||
token, _, err := tokenMgr.GeneratePair("test-user", "testuser")
|
||||
require.NoError(t, err)
|
||||
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws?token=" + token
|
||||
|
||||
return srv, wsURL
|
||||
}
|
||||
@@ -567,3 +574,156 @@ func TestWS_QueryWithTTSDisabled(t *testing.T) {
|
||||
err = conn.ReadJSON(&extra)
|
||||
assert.Error(t, err, "不应有额外消息")
|
||||
}
|
||||
|
||||
// --- 认证测试辅助 ---
|
||||
|
||||
// setupTestServerEx 创建测试服务器,返回 tokenMgr 和 sessionMgr 以便测试控制。
|
||||
func setupTestServerEx(t *testing.T, orch orchestrator.Orchestrator) (*httptest.Server, *auth.TokenManager, *session.MemoryManager) {
|
||||
t.Helper()
|
||||
|
||||
sessionMgr := session.NewMemoryManager(5*time.Minute, 20)
|
||||
t.Cleanup(func() { sessionMgr.Stop() })
|
||||
|
||||
tokenMgr := auth.NewTokenManager("test-secret", 15*time.Minute, 7*24*time.Hour)
|
||||
|
||||
r := gin.New()
|
||||
cfg := &config.Config{
|
||||
App: config.AppConfig{Version: "test"},
|
||||
Server: config.ServerConfig{HeartbeatInterval: 30, HeartbeatTimeout: 60},
|
||||
Session: config.SessionConfig{MaxHistory: 20},
|
||||
}
|
||||
r.GET("/ws", ServeWS(sessionMgr, orch, cfg, tokenMgr))
|
||||
|
||||
srv := httptest.NewServer(r)
|
||||
return srv, tokenMgr, sessionMgr
|
||||
}
|
||||
|
||||
// httpGet 发送 HTTP GET 并返回状态码。
|
||||
func httpGet(t *testing.T, url string) int {
|
||||
t.Helper()
|
||||
resp, err := http.Get(url)
|
||||
require.NoError(t, err)
|
||||
resp.Body.Close()
|
||||
return resp.StatusCode
|
||||
}
|
||||
|
||||
// --- 认证测试用例 ---
|
||||
|
||||
// TestWS_AuthMissingToken 验证无 token 时返回 401。
|
||||
func TestWS_AuthMissingToken(t *testing.T) {
|
||||
srv, _, _ := setupTestServerEx(t, &MockOrchestrator{})
|
||||
defer srv.Close()
|
||||
|
||||
httpURL := srv.URL + "/ws"
|
||||
status := httpGet(t, httpURL)
|
||||
assert.Equal(t, http.StatusUnauthorized, status)
|
||||
}
|
||||
|
||||
// TestWS_AuthInvalidToken 验证无效 token 时返回 401。
|
||||
func TestWS_AuthInvalidToken(t *testing.T) {
|
||||
srv, _, _ := setupTestServerEx(t, &MockOrchestrator{})
|
||||
defer srv.Close()
|
||||
|
||||
httpURL := srv.URL + "/ws?token=invalid-token"
|
||||
status := httpGet(t, httpURL)
|
||||
assert.Equal(t, http.StatusUnauthorized, status)
|
||||
}
|
||||
|
||||
// TestWS_AuthExpiredToken 验证过期 token 时返回 401。
|
||||
func TestWS_AuthExpiredToken(t *testing.T) {
|
||||
// 创建一个 access TTL 极短的 tokenMgr
|
||||
sessionMgr := session.NewMemoryManager(5*time.Minute, 20)
|
||||
defer sessionMgr.Stop()
|
||||
|
||||
tokenMgr := auth.NewTokenManager("test-secret", -1*time.Minute, 7*24*time.Hour) // 已过期
|
||||
|
||||
r := gin.New()
|
||||
cfg := &config.Config{
|
||||
App: config.AppConfig{Version: "test"},
|
||||
Server: config.ServerConfig{HeartbeatInterval: 30, HeartbeatTimeout: 60},
|
||||
Session: config.SessionConfig{MaxHistory: 20},
|
||||
}
|
||||
r.GET("/ws", ServeWS(sessionMgr, &MockOrchestrator{}, cfg, tokenMgr))
|
||||
srv := httptest.NewServer(r)
|
||||
defer srv.Close()
|
||||
|
||||
token, _, err := tokenMgr.GeneratePair("test-user", "testuser")
|
||||
require.NoError(t, err)
|
||||
|
||||
httpURL := srv.URL + "/ws?token=" + token
|
||||
status := httpGet(t, httpURL)
|
||||
assert.Equal(t, http.StatusUnauthorized, status)
|
||||
}
|
||||
|
||||
// TestWS_AuthValidToken 验证有效 token 能成功建立 WS 连接。
|
||||
func TestWS_AuthValidToken(t *testing.T) {
|
||||
srv, tokenMgr, _ := setupTestServerEx(t, &MockOrchestrator{})
|
||||
defer srv.Close()
|
||||
|
||||
token, _, err := tokenMgr.GeneratePair("user-1", "alice")
|
||||
require.NoError(t, err)
|
||||
|
||||
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws?token=" + token
|
||||
conn := connectWS(t, wsURL)
|
||||
|
||||
msg := readJSON(t, conn)
|
||||
assert.Equal(t, "connected", msg["type"])
|
||||
assert.NotEmpty(t, msg["session_id"])
|
||||
}
|
||||
|
||||
// TestWS_AuthConversationIDResume 验证通过 conversation_id 恢复已有对话。
|
||||
func TestWS_AuthConversationIDResume(t *testing.T) {
|
||||
srv, tokenMgr, sessionMgr := setupTestServerEx(t, &MockOrchestrator{})
|
||||
defer srv.Close()
|
||||
|
||||
userID := "user-1"
|
||||
|
||||
// 先创建一个属于该用户的 session
|
||||
ctx := context.Background()
|
||||
sessionID, err := sessionMgr.Create(ctx, userID, models.DefaultConfig())
|
||||
require.NoError(t, err)
|
||||
|
||||
token, _, err := tokenMgr.GeneratePair(userID, "alice")
|
||||
require.NoError(t, err)
|
||||
|
||||
// 带 conversation_id 连接
|
||||
wsURL := "ws" + strings.TrimPrefix(srv.URL, "http") +
|
||||
"/ws?token=" + token + "&conversation_id=" + sessionID
|
||||
conn := connectWS(t, wsURL)
|
||||
|
||||
msg := readJSON(t, conn)
|
||||
assert.Equal(t, "connected", msg["type"])
|
||||
assert.Equal(t, sessionID, msg["session_id"], "应复用已有 session")
|
||||
}
|
||||
|
||||
// TestWS_AuthConversationIDNotFound 验证 conversation_id 不存在时返回 401。
|
||||
func TestWS_AuthConversationIDNotFound(t *testing.T) {
|
||||
srv, tokenMgr, _ := setupTestServerEx(t, &MockOrchestrator{})
|
||||
defer srv.Close()
|
||||
|
||||
token, _, err := tokenMgr.GeneratePair("user-1", "alice")
|
||||
require.NoError(t, err)
|
||||
|
||||
httpURL := srv.URL + "/ws?token=" + token + "&conversation_id=nonexistent-id"
|
||||
status := httpGet(t, httpURL)
|
||||
assert.Equal(t, http.StatusUnauthorized, status)
|
||||
}
|
||||
|
||||
// TestWS_AuthConversationIDOwnership 验证 conversation_id 不属于当前用户时返回 401。
|
||||
func TestWS_AuthConversationIDOwnership(t *testing.T) {
|
||||
srv, tokenMgr, sessionMgr := setupTestServerEx(t, &MockOrchestrator{})
|
||||
defer srv.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
// user-A 创建 session
|
||||
sessionID, err := sessionMgr.Create(ctx, "user-A", models.DefaultConfig())
|
||||
require.NoError(t, err)
|
||||
|
||||
// user-B 尝试连接该 session
|
||||
token, _, err := tokenMgr.GeneratePair("user-B", "bob")
|
||||
require.NoError(t, err)
|
||||
|
||||
httpURL := srv.URL + "/ws?token=" + token + "&conversation_id=" + sessionID
|
||||
status := httpGet(t, httpURL)
|
||||
assert.Equal(t, http.StatusUnauthorized, status, "非 owner 访问应返回 401")
|
||||
}
|
||||
|
||||
1
backend/migrations/002_messages.down.sql
Normal file
1
backend/migrations/002_messages.down.sql
Normal file
@@ -0,0 +1 @@
|
||||
DROP TABLE IF EXISTS messages;
|
||||
17
backend/migrations/002_messages.up.sql
Normal file
17
backend/migrations/002_messages.up.sql
Normal file
@@ -0,0 +1,17 @@
|
||||
-- 消息表
|
||||
CREATE TABLE IF NOT EXISTS messages (
|
||||
id BIGSERIAL PRIMARY KEY,
|
||||
session_id UUID NOT NULL,
|
||||
role VARCHAR(16) NOT NULL, -- "user" | "assistant" | "system"
|
||||
content TEXT NOT NULL,
|
||||
tokens_used INTEGER NOT NULL DEFAULT 0,
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
|
||||
);
|
||||
|
||||
-- 按会话查询消息(分页核心索引)
|
||||
CREATE INDEX IF NOT EXISTS idx_messages_session_id_created_at
|
||||
ON messages(session_id, created_at);
|
||||
|
||||
-- 按会话查询最后一条消息
|
||||
CREATE INDEX IF NOT EXISTS idx_messages_session_id_id_desc
|
||||
ON messages(session_id, id DESC);
|
||||
623
docs/03-接口文档.md
623
docs/03-接口文档.md
@@ -13,16 +13,28 @@
|
||||
|
||||
```
|
||||
浏览器 Go Gateway :8080
|
||||
WebSocket Client <--> /ws (实时对话)
|
||||
HTTP Client --> GET /api/health
|
||||
HTTP Client <--> POST/DELETE /api/sessions
|
||||
WebSocket Client <--> /ws?token=<jwt> (实时对话,需 JWT 认证)
|
||||
HTTP Client --> GET /api/health (健康检查)
|
||||
HTTP Client <--> POST /api/auth/* (注册/登录/刷新/登出)
|
||||
HTTP Client <--> GET/POST/PATCH/DELETE (对话 CRUD)
|
||||
/api/conversations/*
|
||||
HTTP Client <--> GET /api/conversations/:id (历史消息)
|
||||
/messages
|
||||
HTTP Client ~~> POST/DELETE /api/sessions (已废弃,保留兼容)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 一、WebSocket 协议
|
||||
|
||||
连接地址:`ws://localhost:8080/ws`
|
||||
连接地址:`ws://localhost:8080/ws?token=<access_token>&conversation_id=<uuid>`
|
||||
|
||||
| 参数 | 必填 | 说明 |
|
||||
|------|------|------|
|
||||
| `token` | 是 | JWT access_token,缺失或无效时返回 401 |
|
||||
| `conversation_id` | 否 | 恢复已有对话;省略则创建新对话 |
|
||||
|
||||
> 详见"REST API → WebSocket 认证变更"章节。
|
||||
|
||||
### 消息格式约定
|
||||
|
||||
@@ -95,8 +107,9 @@ interface PingMessage {
|
||||
```typescript
|
||||
interface ConnectedMessage {
|
||||
type: "connected";
|
||||
session_id: string; // 服务端生成的会话 ID
|
||||
server_version: string; // 服务端版本号,如 "0.1.0"
|
||||
session_id: string; // 服务端生成的会话 ID
|
||||
conversation_id: string; // 同 session_id,便于前端统一使用
|
||||
server_version: string; // 服务端版本号,如 "0.1.0"
|
||||
}
|
||||
```
|
||||
|
||||
@@ -265,13 +278,446 @@ Client Server
|
||||
|
||||
## 二、REST API
|
||||
|
||||
### 通用约定
|
||||
|
||||
#### 认证方式
|
||||
|
||||
需要认证的接口在请求头携带 JWT access token:
|
||||
|
||||
```
|
||||
Authorization: Bearer <access_token>
|
||||
```
|
||||
|
||||
未认证或 token 过期时返回 `401 Unauthorized`。
|
||||
|
||||
#### 错误响应格式
|
||||
|
||||
所有错误响应统一结构:
|
||||
|
||||
```typescript
|
||||
interface ApiError {
|
||||
code: string; // 机器可读错误码
|
||||
message: string; // 人类可读描述
|
||||
}
|
||||
```
|
||||
|
||||
示例:
|
||||
|
||||
```json
|
||||
{
|
||||
"code": "USERNAME_TAKEN",
|
||||
"message": "username already taken"
|
||||
}
|
||||
```
|
||||
|
||||
#### 输入校验规则
|
||||
|
||||
| 字段 | 规则 |
|
||||
|------|------|
|
||||
| `username` | 3-64 字符,仅允许字母、数字、下划线 |
|
||||
| `password` | 8-72 字符 |
|
||||
|
||||
---
|
||||
|
||||
### 认证接口(`/api/auth`)
|
||||
|
||||
#### 注册
|
||||
|
||||
```
|
||||
POST /api/auth/register
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
**请求体**:
|
||||
|
||||
```typescript
|
||||
interface RegisterRequest {
|
||||
username: string; // 3-64 字符
|
||||
password: string; // 8-72 字符
|
||||
}
|
||||
```
|
||||
|
||||
**成功响应** `201 Created`:
|
||||
|
||||
```typescript
|
||||
interface AuthResponse {
|
||||
user: {
|
||||
id: string; // UUID
|
||||
username: string;
|
||||
created_at: string; // ISO 8601
|
||||
};
|
||||
access_token: string; // JWT,15 分钟有效
|
||||
refresh_token: string; // JWT,7 天有效
|
||||
}
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"user": {
|
||||
"id": "550e8400-e29b-41d4-a716-446655440000",
|
||||
"username": "alice",
|
||||
"created_at": "2026-06-14T10:00:00Z"
|
||||
},
|
||||
"access_token": "eyJhbGciOiJIUzI1NiIs...",
|
||||
"refresh_token": "eyJhbGciOiJIUzI1NiIs..."
|
||||
}
|
||||
```
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 400 | `INVALID_INPUT` | 用户名/密码不符合校验规则 |
|
||||
| 409 | `USERNAME_TAKEN` | 用户名已存在 |
|
||||
|
||||
#### 登录
|
||||
|
||||
```
|
||||
POST /api/auth/login
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
**请求体**:
|
||||
|
||||
```typescript
|
||||
interface LoginRequest {
|
||||
username: string;
|
||||
password: string;
|
||||
}
|
||||
```
|
||||
|
||||
**成功响应** `200 OK`:同 `AuthResponse` 结构。
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 400 | `INVALID_INPUT` | 请求参数缺失或格式错误 |
|
||||
| 401 | `INVALID_CREDENTIALS` | 用户名或密码错误 |
|
||||
|
||||
#### 刷新 Token
|
||||
|
||||
```
|
||||
POST /api/auth/refresh
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
**请求体**:
|
||||
|
||||
```typescript
|
||||
interface RefreshRequest {
|
||||
refresh_token: string; // 之前签发的 refresh_token
|
||||
}
|
||||
```
|
||||
|
||||
**成功响应** `200 OK`:同 `AuthResponse` 结构(返回新的 access_token + refresh_token,旧 refresh_token 失效——Token 轮转)。
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 401 | `INVALID_TOKEN` | refresh_token 无效或已过期 |
|
||||
|
||||
#### 登出
|
||||
|
||||
```
|
||||
POST /api/auth/logout
|
||||
Content-Type: application/json
|
||||
Authorization: Bearer <access_token>
|
||||
```
|
||||
|
||||
**请求体**:
|
||||
|
||||
```typescript
|
||||
interface LogoutRequest {
|
||||
refresh_token: string; // 要废弃的 refresh_token
|
||||
}
|
||||
```
|
||||
|
||||
**成功响应** `204 No Content`(无响应体)。
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 401 | `INVALID_TOKEN` | access_token 无效或已过期 |
|
||||
|
||||
---
|
||||
|
||||
### 对话接口(`/api/conversations`)
|
||||
|
||||
> 以下所有接口均需认证(`Authorization: Bearer <access_token>`),省略不重复标注。
|
||||
|
||||
#### 对话列表
|
||||
|
||||
```
|
||||
GET /api/conversations?page=1&size=20
|
||||
```
|
||||
|
||||
**查询参数**:
|
||||
|
||||
| 参数 | 类型 | 默认值 | 说明 |
|
||||
|------|------|--------|------|
|
||||
| `page` | int | 1 | 页码,从 1 开始 |
|
||||
| `size` | int | 20 | 每页条数,最大 50 |
|
||||
|
||||
**成功响应** `200 OK`:
|
||||
|
||||
```typescript
|
||||
interface ConversationListResponse {
|
||||
conversations: ConversationSummary[];
|
||||
total: number; // 总条数
|
||||
page: number;
|
||||
size: number;
|
||||
}
|
||||
|
||||
interface ConversationSummary {
|
||||
id: string; // 对话 ID(即 session_id)
|
||||
title: string; // 对话标题(首条消息前 20 字)
|
||||
last_message: string; // 最后一条消息内容预览
|
||||
message_count: number; // 消息总数
|
||||
updated_at: string; // ISO 8601,最后活跃时间
|
||||
}
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"conversations": [
|
||||
{
|
||||
"id": "550e8400-e29b-41d4-a716-446655440000",
|
||||
"title": "这是一朵红色的玫瑰…",
|
||||
"last_message": "它看起来很美丽。",
|
||||
"message_count": 4,
|
||||
"updated_at": "2026-06-14T10:05:30Z"
|
||||
}
|
||||
],
|
||||
"total": 1,
|
||||
"page": 1,
|
||||
"size": 20
|
||||
}
|
||||
```
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 401 | `INVALID_TOKEN` | 未认证或 token 过期 |
|
||||
|
||||
#### 创建对话
|
||||
|
||||
```
|
||||
POST /api/conversations
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
**请求体**(可选,全部有默认值):
|
||||
|
||||
```typescript
|
||||
interface CreateConversationRequest {
|
||||
config?: {
|
||||
tts_enabled?: boolean; // 默认 true
|
||||
detail_level?: "low" | "high"; // 默认 "low"
|
||||
language?: string; // 默认 "zh-CN"
|
||||
};
|
||||
}
|
||||
```
|
||||
|
||||
**成功响应** `201 Created`:
|
||||
|
||||
```typescript
|
||||
interface ConversationDetail {
|
||||
id: string;
|
||||
title: string;
|
||||
config: {
|
||||
tts_enabled: boolean;
|
||||
detail_level: "low" | "high";
|
||||
language: string;
|
||||
};
|
||||
created_at: string; // ISO 8601
|
||||
}
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"id": "660e8400-e29b-41d4-a716-446655440001",
|
||||
"title": "新对话",
|
||||
"config": {
|
||||
"tts_enabled": true,
|
||||
"detail_level": "low",
|
||||
"language": "zh-CN"
|
||||
},
|
||||
"created_at": "2026-06-14T11:00:00Z"
|
||||
}
|
||||
```
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 401 | `INVALID_TOKEN` | 未认证或 token 过期 |
|
||||
|
||||
#### 获取对话详情
|
||||
|
||||
```
|
||||
GET /api/conversations/:id
|
||||
```
|
||||
|
||||
**成功响应** `200 OK`:同 `ConversationDetail` 结构。
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 401 | `INVALID_TOKEN` | 未认证或 token 过期 |
|
||||
| 404 | `SESSION_NOT_FOUND` | 对话不存在或不属于当前用户 |
|
||||
|
||||
#### 更新对话标题
|
||||
|
||||
```
|
||||
PATCH /api/conversations/:id
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
**请求体**:
|
||||
|
||||
```typescript
|
||||
interface UpdateTitleRequest {
|
||||
title: string; // 1-100 字符
|
||||
}
|
||||
```
|
||||
|
||||
**成功响应** `200 OK`:
|
||||
|
||||
```json
|
||||
{
|
||||
"id": "550e8400-e29b-41d4-a716-446655440000",
|
||||
"title": "新的自定义标题"
|
||||
}
|
||||
```
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 400 | `INVALID_INPUT` | title 为空或超长 |
|
||||
| 401 | `INVALID_TOKEN` | 未认证或 token 过期 |
|
||||
| 404 | `SESSION_NOT_FOUND` | 对话不存在或不属于当前用户 |
|
||||
|
||||
#### 删除对话
|
||||
|
||||
```
|
||||
DELETE /api/conversations/:id
|
||||
```
|
||||
|
||||
**成功响应** `204 No Content`(无响应体)。
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 401 | `INVALID_TOKEN` | 未认证或 token 过期 |
|
||||
| 404 | `SESSION_NOT_FOUND` | 对话不存在或不属于当前用户 |
|
||||
|
||||
#### 获取对话消息
|
||||
|
||||
```
|
||||
GET /api/conversations/:id/messages?limit=50&before=<message_id>
|
||||
```
|
||||
|
||||
**查询参数**:
|
||||
|
||||
| 参数 | 类型 | 默认值 | 说明 |
|
||||
|------|------|--------|------|
|
||||
| `limit` | int | 50 | 返回条数,最大 100 |
|
||||
| `before` | int64 | — | 游标分页:返回此 message_id 之前的消息(不含),用于加载更多 |
|
||||
|
||||
**成功响应** `200 OK`:
|
||||
|
||||
```typescript
|
||||
interface MessagesResponse {
|
||||
messages: StoredMessage[];
|
||||
has_more: boolean; // 是否还有更早的消息
|
||||
}
|
||||
|
||||
interface StoredMessage {
|
||||
id: number; // 自增 ID,用于游标分页
|
||||
role: "user" | "assistant";
|
||||
content: string;
|
||||
tokens_used: number; // 该条消息消耗的 token 数
|
||||
created_at: string; // ISO 8601
|
||||
}
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"messages": [
|
||||
{
|
||||
"id": 1001,
|
||||
"role": "user",
|
||||
"content": "这是什么花?",
|
||||
"tokens_used": 0,
|
||||
"created_at": "2026-06-14T10:01:00Z"
|
||||
},
|
||||
{
|
||||
"id": 1002,
|
||||
"role": "assistant",
|
||||
"content": "这是一朵红色的玫瑰。",
|
||||
"tokens_used": 42,
|
||||
"created_at": "2026-06-14T10:01:02Z"
|
||||
}
|
||||
],
|
||||
"has_more": false
|
||||
}
|
||||
```
|
||||
|
||||
**分页用法**:首次请求不带 `before`,获取最新消息。滚动到顶部时,取当前列表最小的 `id` 作为 `before` 参数请求更早的消息。
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 401 | `INVALID_TOKEN` | 未认证或 token 过期 |
|
||||
| 404 | `SESSION_NOT_FOUND` | 对话不存在或不属于当前用户 |
|
||||
|
||||
---
|
||||
|
||||
### WebSocket 认证变更
|
||||
|
||||
连接地址变更为带 token 的查询参数:
|
||||
|
||||
```
|
||||
ws://localhost:8080/ws?token=<access_token>&conversation_id=<uuid>
|
||||
```
|
||||
|
||||
| 参数 | 必填 | 说明 |
|
||||
|------|------|------|
|
||||
| `token` | 是 | JWT access_token |
|
||||
| `conversation_id` | 否 | 恢复已有对话;省略则创建新对话 |
|
||||
|
||||
**认证失败响应**(HTTP 升级前返回):
|
||||
|
||||
| 状态码 | 场景 |
|
||||
|--------|------|
|
||||
| 401 | token 缺失、无效或已过期 |
|
||||
|
||||
**conversation_id 校验失败**:
|
||||
|
||||
| 场景 | 处理 |
|
||||
|------|------|
|
||||
| 对话不存在 | 返回 401,`{"error": "SESSION_NOT_FOUND"}` |
|
||||
| 对话不属于当前用户 | 返回 401,`{"error": "SESSION_NOT_FOUND"}`(与不存在相同,避免信息泄露) |
|
||||
|
||||
---
|
||||
|
||||
### 健康检查
|
||||
|
||||
```
|
||||
GET /api/health
|
||||
```
|
||||
|
||||
响应:
|
||||
无需认证。
|
||||
|
||||
**成功响应** `200 OK`:
|
||||
|
||||
```json
|
||||
{
|
||||
@@ -282,43 +728,21 @@ GET /api/health
|
||||
}
|
||||
```
|
||||
|
||||
### 创建会话(可选,MVP 自动创建)
|
||||
---
|
||||
|
||||
### ~~旧会话接口~~(已废弃)
|
||||
|
||||
> 以下端点已废弃,保留仅为向后兼容。新代码应使用 `/api/conversations` 系列接口。
|
||||
|
||||
```
|
||||
POST /api/sessions
|
||||
Content-Type: application/json
|
||||
|
||||
{
|
||||
"config": {
|
||||
"tts_enabled": true,
|
||||
"detail_level": "low",
|
||||
"language": "zh-CN"
|
||||
}
|
||||
}
|
||||
POST /api/sessions → 改用 POST /api/conversations
|
||||
DELETE /api/sessions/{id} → 改用 DELETE /api/conversations/{id}
|
||||
```
|
||||
|
||||
响应:
|
||||
|
||||
```json
|
||||
{
|
||||
"session_id": "550e8400-e29b-41d4-a716-446655440000",
|
||||
"created_at": "2026-06-12T15:41:00Z"
|
||||
}
|
||||
```
|
||||
|
||||
### 销毁会话
|
||||
|
||||
```
|
||||
DELETE /api/sessions/{session_id}
|
||||
```
|
||||
|
||||
响应:`204 No Content`
|
||||
|
||||
### 预留端点(暂不实现)
|
||||
|
||||
| 端点 | 方法 | 用途 |
|
||||
|------|------|------|
|
||||
| `/api/sessions/{id}/messages` | GET | 查询对话历史 |
|
||||
| `/api/usage` | GET | 查询用量统计 |
|
||||
| `/api/users/{id}/preferences` | GET/PUT | 用户偏好管理 |
|
||||
|
||||
@@ -689,6 +1113,7 @@ type Config struct {
|
||||
Redis RedisConfig `mapstructure:"redis"`
|
||||
AI AIConfig `mapstructure:"ai"`
|
||||
Storage StorageConfig `mapstructure:"storage"`
|
||||
Auth AuthConfig `mapstructure:"auth"`
|
||||
Log LogConfig `mapstructure:"log"`
|
||||
}
|
||||
|
||||
@@ -746,6 +1171,12 @@ type StorageConfig struct {
|
||||
DSN string `mapstructure:"dsn"` // PostgreSQL 连接串,driver=postgres 时必填
|
||||
}
|
||||
|
||||
type AuthConfig struct {
|
||||
JWTSecret string `mapstructure:"jwt_secret"` // 必须通过 CAMTALK_AUTH_JWT_SECRET 设置
|
||||
AccessTTL int `mapstructure:"access_ttl"` // 分钟,默认 15
|
||||
RefreshTTL int `mapstructure:"refresh_ttl"` // 分钟,默认 10080(7 天)
|
||||
}
|
||||
|
||||
type LogConfig struct {
|
||||
Level string `mapstructure:"level"` // "debug" | "info" | "warn" | "error",默认 "info"
|
||||
Format string `mapstructure:"format"` // "json" | "console",生产用 json
|
||||
@@ -791,6 +1222,10 @@ ai:
|
||||
storage:
|
||||
driver: memory
|
||||
|
||||
auth:
|
||||
access_ttl: 15 # access token 有效期(分钟)
|
||||
refresh_ttl: 10080 # refresh token 有效期(分钟,7 天)
|
||||
|
||||
log:
|
||||
level: info
|
||||
format: console
|
||||
@@ -811,6 +1246,9 @@ Viper 自动将配置项映射为环境变量,规则:**前缀 `CAMTALK_` +
|
||||
| `ai.llm.model` | `CAMTALK_AI_LLM_MODEL` | `gpt-4o` |
|
||||
| `storage.driver` | `CAMTALK_STORAGE_DRIVER` | `postgres` |
|
||||
| `storage.dsn` | `CAMTALK_STORAGE_DSN` | — |
|
||||
| `auth.jwt_secret` | `CAMTALK_AUTH_JWT_SECRET` | —(必填,仅环境变量) |
|
||||
| `auth.access_ttl` | `CAMTALK_AUTH_ACCESS_TTL` | `15` |
|
||||
| `auth.refresh_ttl` | `CAMTALK_AUTH_REFRESH_TTL` | `10080` |
|
||||
| `app.env` | `CAMTALK_APP_ENV` | `prod` |
|
||||
| `log.level` | `CAMTALK_LOG_LEVEL` | `warn` |
|
||||
| `log.format` | `CAMTALK_LOG_FORMAT` | `json` |
|
||||
@@ -892,6 +1330,7 @@ CAMTALK_AI_STT_API_KEY=xxx \
|
||||
CAMTALK_AI_TTS_API_KEY=xxx \
|
||||
CAMTALK_STORAGE_DRIVER=postgres \
|
||||
CAMTALK_STORAGE_DSN="postgres://user:pass@db:5432/camtalk?sslmode=disable" \
|
||||
CAMTALK_AUTH_JWT_SECRET="$(openssl rand -hex 32)" \
|
||||
CAMTALK_LOG_LEVEL=warn \
|
||||
CAMTALK_LOG_FORMAT=json \
|
||||
./bin/camtalk
|
||||
@@ -910,7 +1349,10 @@ CAMTALK_LOG_FORMAT=json \
|
||||
|
||||
type Session struct {
|
||||
ID string `json:"session_id"`
|
||||
UserID string `json:"user_id"` // 关联用户,空串表示匿名
|
||||
Title string `json:"title"` // 对话标题
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
Config SessionConfig `json:"config"`
|
||||
}
|
||||
|
||||
@@ -932,6 +1374,33 @@ type Message struct {
|
||||
Role string `json:"role"` // "user" | "assistant"
|
||||
Content string `json:"content"`
|
||||
}
|
||||
|
||||
// ---- 用户模块 ----
|
||||
|
||||
type User struct {
|
||||
ID string `json:"id"`
|
||||
Username string `json:"username"`
|
||||
PasswordHash string `json:"-"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
type ConversationSummary struct {
|
||||
ID string `json:"id"`
|
||||
Title string `json:"title"`
|
||||
LastMessage string `json:"last_message"`
|
||||
MessageCount int `json:"message_count"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
}
|
||||
|
||||
type StoredMessage struct {
|
||||
ID int64 `json:"id"`
|
||||
SessionID string `json:"-"`
|
||||
Role string `json:"role"`
|
||||
Content string `json:"content"`
|
||||
TokensUsed int `json:"tokens_used"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
```
|
||||
|
||||
### TypeScript 前端模型
|
||||
@@ -957,6 +1426,60 @@ interface ChatMessage {
|
||||
tokensUsed?: number;
|
||||
}
|
||||
|
||||
// ---- 用户模块 ----
|
||||
|
||||
interface AuthTokens {
|
||||
accessToken: string;
|
||||
refreshToken: string;
|
||||
}
|
||||
|
||||
interface User {
|
||||
id: string; // UUID
|
||||
username: string;
|
||||
created_at: string; // ISO 8601
|
||||
}
|
||||
|
||||
interface AuthResponse {
|
||||
user: User;
|
||||
access_token: string;
|
||||
refresh_token: string;
|
||||
}
|
||||
|
||||
interface ConversationSummary {
|
||||
id: string;
|
||||
title: string;
|
||||
last_message: string;
|
||||
message_count: number;
|
||||
updated_at: string;
|
||||
}
|
||||
|
||||
interface ConversationListResponse {
|
||||
conversations: ConversationSummary[];
|
||||
total: number;
|
||||
page: number;
|
||||
size: number;
|
||||
}
|
||||
|
||||
interface ConversationDetail {
|
||||
id: string;
|
||||
title: string;
|
||||
config: SessionConfig;
|
||||
created_at: string;
|
||||
}
|
||||
|
||||
interface StoredMessage {
|
||||
id: number;
|
||||
role: "user" | "assistant";
|
||||
content: string;
|
||||
tokens_used: number;
|
||||
created_at: string;
|
||||
}
|
||||
|
||||
interface MessagesResponse {
|
||||
messages: StoredMessage[];
|
||||
has_more: boolean;
|
||||
}
|
||||
|
||||
// WebSocket 消息联合类型
|
||||
type ServerMessage =
|
||||
| ConnectedMessage
|
||||
@@ -1054,18 +1577,22 @@ func NewApp(cfg *Config) *App {
|
||||
|
||||
## 九、错误码
|
||||
|
||||
| 错误码 | 含义 | 客户端处理建议 |
|
||||
|--------|------|--------------|
|
||||
| `INVALID_MESSAGE` | 消息格式不合法 | 检查 JSON 结构,不重试 |
|
||||
| `SESSION_NOT_FOUND` | 会话不存在或已过期 | 重新建立 WebSocket 连接 |
|
||||
| `RATE_LIMITED` | 请求频率超限 | 延迟后重试,提示用户稍等 |
|
||||
| `IMAGE_TOO_LARGE` | 图像超过 4MB 限制 | 降低分辨率或压缩质量 |
|
||||
| `AUDIO_TOO_SHORT` | 音频片段 < 250ms | 忽略,等待下次语音输入 |
|
||||
| `LLM_TIMEOUT` | LLM 推理超时(>10s) | 提示用户重试 |
|
||||
| `LLM_ERROR` | LLM 服务异常 | 提示用户重试,服务端记录日志 |
|
||||
| `STT_ERROR` | 语音识别失败 | 回退到纯文本输入模式 |
|
||||
| `TTS_ERROR` | 语音合成失败 | 静默回退到纯文本回复 |
|
||||
| `INTERNAL_ERROR` | 服务端内部错误 | 提示用户重试 |
|
||||
| 错误码 | HTTP 状态码 | 含义 | 客户端处理建议 |
|
||||
|--------|-----------|------|--------------|
|
||||
| `INVALID_MESSAGE` | — | 消息格式不合法(WS) | 检查 JSON 结构,不重试 |
|
||||
| `SESSION_NOT_FOUND` | 404 | 会话/对话不存在或已过期 | 重新建立连接或刷新列表 |
|
||||
| `RATE_LIMITED` | 429 | 请求频率超限 | 延迟后重试,提示用户稍等 |
|
||||
| `IMAGE_TOO_LARGE` | — | 图像超过 4MB 限制(WS) | 降低分辨率或压缩质量 |
|
||||
| `AUDIO_TOO_SHORT` | — | 音频片段 < 250ms(WS) | 忽略,等待下次语音输入 |
|
||||
| `LLM_TIMEOUT` | — | LLM 推理超时 >10s(WS) | 提示用户重试 |
|
||||
| `LLM_ERROR` | — | LLM 服务异常(WS) | 提示用户重试,服务端记录日志 |
|
||||
| `STT_ERROR` | — | 语音识别失败(WS) | 回退到纯文本输入模式 |
|
||||
| `TTS_ERROR` | — | 语音合成失败(WS) | 静默回退到纯文本回复 |
|
||||
| `INTERNAL_ERROR` | 500 | 服务端内部错误 | 提示用户重试 |
|
||||
| `USERNAME_TAKEN` | 409 | 用户名已被注册 | 提示换一个用户名 |
|
||||
| `INVALID_CREDENTIALS` | 401 | 用户名或密码错误 | 提示检查输入 |
|
||||
| `INVALID_TOKEN` | 401 | JWT 无效或已过期 | 尝试 refresh,失败则重新登录 |
|
||||
| `INVALID_INPUT` | 400 | 请求参数校验失败 | 检查字段规则后重试 |
|
||||
|
||||
## 十、连接管理
|
||||
|
||||
|
||||
@@ -713,6 +713,588 @@ func main() {
|
||||
|
||||
---
|
||||
|
||||
## 前端 API 接口参考
|
||||
|
||||
本章节为前端开发者提供完整的 REST API 契约。所有接口以 JSON 通信,基地址与 WebSocket 同源(开发环境 `http://localhost:8080`,生产环境通过 Nginx 反代)。
|
||||
|
||||
### 通用约定
|
||||
|
||||
#### 认证方式
|
||||
|
||||
需要认证的接口在请求头携带 JWT access token:
|
||||
|
||||
```
|
||||
Authorization: Bearer <access_token>
|
||||
```
|
||||
|
||||
未认证或 token 过期时返回 `401 Unauthorized`。
|
||||
|
||||
#### 错误响应格式
|
||||
|
||||
所有错误响应统一结构:
|
||||
|
||||
```typescript
|
||||
interface ApiError {
|
||||
code: string; // 机器可读错误码
|
||||
message: string; // 人类可读描述
|
||||
}
|
||||
```
|
||||
|
||||
示例:
|
||||
|
||||
```json
|
||||
{
|
||||
"code": "USERNAME_TAKEN",
|
||||
"message": "username already taken"
|
||||
}
|
||||
```
|
||||
|
||||
#### 新增错误码
|
||||
|
||||
| 错误码 | HTTP 状态码 | 含义 |
|
||||
|--------|-----------|------|
|
||||
| `USERNAME_TAKEN` | 409 | 用户名已被注册 |
|
||||
| `INVALID_CREDENTIALS` | 401 | 用户名或密码错误 |
|
||||
| `INVALID_TOKEN` | 401 | JWT 无效或已过期 |
|
||||
| `INVALID_INPUT` | 400 | 请求参数校验失败 |
|
||||
| `SESSION_NOT_FOUND` | 404 | 对话不存在或无权访问 |
|
||||
|
||||
#### 输入校验规则
|
||||
|
||||
| 字段 | 规则 |
|
||||
|------|------|
|
||||
| `username` | 3-64 字符,仅允许字母、数字、下划线 |
|
||||
| `password` | 8-72 字符 |
|
||||
|
||||
---
|
||||
|
||||
### 一、认证接口(`/api/auth`)
|
||||
|
||||
#### 1.1 注册
|
||||
|
||||
```
|
||||
POST /api/auth/register
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
**请求体**:
|
||||
|
||||
```typescript
|
||||
interface RegisterRequest {
|
||||
username: string; // 3-64 字符
|
||||
password: string; // 8-72 字符
|
||||
}
|
||||
```
|
||||
|
||||
**成功响应** `201 Created`:
|
||||
|
||||
```typescript
|
||||
interface AuthResponse {
|
||||
user: {
|
||||
id: string; // UUID
|
||||
username: string;
|
||||
created_at: string; // ISO 8601
|
||||
};
|
||||
access_token: string; // JWT,15 分钟有效
|
||||
refresh_token: string; // JWT,7 天有效
|
||||
}
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"user": {
|
||||
"id": "550e8400-e29b-41d4-a716-446655440000",
|
||||
"username": "alice",
|
||||
"created_at": "2026-06-14T10:00:00Z"
|
||||
},
|
||||
"access_token": "eyJhbGciOiJIUzI1NiIs...",
|
||||
"refresh_token": "eyJhbGciOiJIUzI1NiIs..."
|
||||
}
|
||||
```
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 400 | `INVALID_INPUT` | 用户名/密码不符合校验规则 |
|
||||
| 409 | `USERNAME_TAKEN` | 用户名已存在 |
|
||||
|
||||
---
|
||||
|
||||
#### 1.2 登录
|
||||
|
||||
```
|
||||
POST /api/auth/login
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
**请求体**:
|
||||
|
||||
```typescript
|
||||
interface LoginRequest {
|
||||
username: string;
|
||||
password: string;
|
||||
}
|
||||
```
|
||||
|
||||
**成功响应** `200 OK`:同 `AuthResponse` 结构。
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 400 | `INVALID_INPUT` | 请求参数缺失或格式错误 |
|
||||
| 401 | `INVALID_CREDENTIALS` | 用户名或密码错误 |
|
||||
|
||||
---
|
||||
|
||||
#### 1.3 刷新 Token
|
||||
|
||||
```
|
||||
POST /api/auth/refresh
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
**请求体**:
|
||||
|
||||
```typescript
|
||||
interface RefreshRequest {
|
||||
refresh_token: string; // 之前签发的 refresh_token
|
||||
}
|
||||
```
|
||||
|
||||
**成功响应** `200 OK`:同 `AuthResponse` 结构(返回新的 access_token + refresh_token,旧 refresh_token 失效——Token 轮转)。
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 401 | `INVALID_TOKEN` | refresh_token 无效或已过期 |
|
||||
|
||||
---
|
||||
|
||||
#### 1.4 登出
|
||||
|
||||
```
|
||||
POST /api/auth/logout
|
||||
Content-Type: application/json
|
||||
Authorization: Bearer <access_token>
|
||||
```
|
||||
|
||||
**请求体**:
|
||||
|
||||
```typescript
|
||||
interface LogoutRequest {
|
||||
refresh_token: string; // 要废弃的 refresh_token
|
||||
}
|
||||
```
|
||||
|
||||
**成功响应** `204 No Content`(无响应体)。
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 401 | `INVALID_TOKEN` | access_token 无效或已过期 |
|
||||
|
||||
---
|
||||
|
||||
### 二、对话接口(`/api/conversations`)
|
||||
|
||||
> 以下所有接口均需认证(`Authorization: Bearer <access_token>`),省略不重复标注。
|
||||
|
||||
#### 2.1 对话列表
|
||||
|
||||
```
|
||||
GET /api/conversations?page=1&size=20
|
||||
```
|
||||
|
||||
**查询参数**:
|
||||
|
||||
| 参数 | 类型 | 默认值 | 说明 |
|
||||
|------|------|--------|------|
|
||||
| `page` | int | 1 | 页码,从 1 开始 |
|
||||
| `size` | int | 20 | 每页条数,最大 50 |
|
||||
|
||||
**成功响应** `200 OK`:
|
||||
|
||||
```typescript
|
||||
interface ConversationListResponse {
|
||||
conversations: ConversationSummary[];
|
||||
total: number; // 总条数
|
||||
page: number;
|
||||
size: number;
|
||||
}
|
||||
|
||||
interface ConversationSummary {
|
||||
id: string; // 对话 ID(即 session_id)
|
||||
title: string; // 对话标题(首条消息前 20 字)
|
||||
last_message: string; // 最后一条消息内容预览
|
||||
message_count: number; // 消息总数
|
||||
updated_at: string; // ISO 8601,最后活跃时间
|
||||
}
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"conversations": [
|
||||
{
|
||||
"id": "550e8400-e29b-41d4-a716-446655440000",
|
||||
"title": "这是一朵红色的玫瑰…",
|
||||
"last_message": "它看起来很美丽。",
|
||||
"message_count": 4,
|
||||
"updated_at": "2026-06-14T10:05:30Z"
|
||||
}
|
||||
],
|
||||
"total": 1,
|
||||
"page": 1,
|
||||
"size": 20
|
||||
}
|
||||
```
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 401 | `INVALID_TOKEN` | 未认证或 token 过期 |
|
||||
|
||||
---
|
||||
|
||||
#### 2.2 创建对话
|
||||
|
||||
```
|
||||
POST /api/conversations
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
**请求体**(可选,全部有默认值):
|
||||
|
||||
```typescript
|
||||
interface CreateConversationRequest {
|
||||
config?: {
|
||||
tts_enabled?: boolean; // 默认 true
|
||||
detail_level?: "low" | "high"; // 默认 "low"
|
||||
language?: string; // 默认 "zh-CN"
|
||||
};
|
||||
}
|
||||
```
|
||||
|
||||
**成功响应** `201 Created`:
|
||||
|
||||
```typescript
|
||||
interface ConversationDetail {
|
||||
id: string;
|
||||
title: string;
|
||||
config: {
|
||||
tts_enabled: boolean;
|
||||
detail_level: "low" | "high";
|
||||
language: string;
|
||||
};
|
||||
created_at: string; // ISO 8601
|
||||
}
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"id": "660e8400-e29b-41d4-a716-446655440001",
|
||||
"title": "新对话",
|
||||
"config": {
|
||||
"tts_enabled": true,
|
||||
"detail_level": "low",
|
||||
"language": "zh-CN"
|
||||
},
|
||||
"created_at": "2026-06-14T11:00:00Z"
|
||||
}
|
||||
```
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 401 | `INVALID_TOKEN` | 未认证或 token 过期 |
|
||||
|
||||
---
|
||||
|
||||
#### 2.3 获取对话详情
|
||||
|
||||
```
|
||||
GET /api/conversations/:id
|
||||
```
|
||||
|
||||
**成功响应** `200 OK`:同 `ConversationDetail` 结构。
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 401 | `INVALID_TOKEN` | 未认证或 token 过期 |
|
||||
| 404 | `SESSION_NOT_FOUND` | 对话不存在或不属于当前用户 |
|
||||
|
||||
---
|
||||
|
||||
#### 2.4 更新对话标题
|
||||
|
||||
```
|
||||
PATCH /api/conversations/:id
|
||||
Content-Type: application/json
|
||||
```
|
||||
|
||||
**请求体**:
|
||||
|
||||
```typescript
|
||||
interface UpdateTitleRequest {
|
||||
title: string; // 1-100 字符
|
||||
}
|
||||
```
|
||||
|
||||
**成功响应** `200 OK`:
|
||||
|
||||
```json
|
||||
{
|
||||
"id": "550e8400-e29b-41d4-a716-446655440000",
|
||||
"title": "新的自定义标题"
|
||||
}
|
||||
```
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 400 | `INVALID_INPUT` | title 为空或超长 |
|
||||
| 401 | `INVALID_TOKEN` | 未认证或 token 过期 |
|
||||
| 404 | `SESSION_NOT_FOUND` | 对话不存在或不属于当前用户 |
|
||||
|
||||
---
|
||||
|
||||
#### 2.5 删除对话
|
||||
|
||||
```
|
||||
DELETE /api/conversations/:id
|
||||
```
|
||||
|
||||
**成功响应** `204 No Content`(无响应体)。
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 401 | `INVALID_TOKEN` | 未认证或 token 过期 |
|
||||
| 404 | `SESSION_NOT_FOUND` | 对话不存在或不属于当前用户 |
|
||||
|
||||
---
|
||||
|
||||
#### 2.6 获取对话消息
|
||||
|
||||
```
|
||||
GET /api/conversations/:id/messages?limit=50&before=<message_id>
|
||||
```
|
||||
|
||||
**查询参数**:
|
||||
|
||||
| 参数 | 类型 | 默认值 | 说明 |
|
||||
|------|------|--------|------|
|
||||
| `limit` | int | 50 | 返回条数,最大 100 |
|
||||
| `before` | int64 | — | 游标分页:返回此 message_id 之前的消息(不含),用于加载更多 |
|
||||
|
||||
**成功响应** `200 OK`:
|
||||
|
||||
```typescript
|
||||
interface MessagesResponse {
|
||||
messages: StoredMessage[];
|
||||
has_more: boolean; // 是否还有更早的消息
|
||||
}
|
||||
|
||||
interface StoredMessage {
|
||||
id: number; // 自增 ID,用于游标分页
|
||||
role: "user" | "assistant";
|
||||
content: string;
|
||||
tokens_used: number; // 该条消息消耗的 token 数
|
||||
created_at: string; // ISO 8601
|
||||
}
|
||||
```
|
||||
|
||||
```json
|
||||
{
|
||||
"messages": [
|
||||
{
|
||||
"id": 1001,
|
||||
"role": "user",
|
||||
"content": "这是什么花?",
|
||||
"tokens_used": 0,
|
||||
"created_at": "2026-06-14T10:01:00Z"
|
||||
},
|
||||
{
|
||||
"id": 1002,
|
||||
"role": "assistant",
|
||||
"content": "这是一朵红色的玫瑰。",
|
||||
"tokens_used": 42,
|
||||
"created_at": "2026-06-14T10:01:02Z"
|
||||
}
|
||||
],
|
||||
"has_more": false
|
||||
}
|
||||
```
|
||||
|
||||
**分页用法**:首次请求不带 `before`,获取最新消息。滚动到顶部时,取当前列表最小的 `id` 作为 `before` 参数请求更早的消息。
|
||||
|
||||
**错误响应**:
|
||||
|
||||
| 状态码 | code | 场景 |
|
||||
|--------|------|------|
|
||||
| 401 | `INVALID_TOKEN` | 未认证或 token 过期 |
|
||||
| 404 | `SESSION_NOT_FOUND` | 对话不存在或不属于当前用户 |
|
||||
|
||||
---
|
||||
|
||||
### 三、WebSocket 认证变更
|
||||
|
||||
连接地址变更为带 token 的查询参数:
|
||||
|
||||
```
|
||||
ws://localhost:8080/ws?token=<access_token>&conversation_id=<uuid>
|
||||
```
|
||||
|
||||
| 参数 | 必填 | 说明 |
|
||||
|------|------|------|
|
||||
| `token` | 是 | JWT access_token |
|
||||
| `conversation_id` | 否 | 恢复已有对话;省略则创建新对话 |
|
||||
|
||||
**认证失败响应**(HTTP 升级前返回):
|
||||
|
||||
| 状态码 | 场景 |
|
||||
|--------|------|
|
||||
| 401 | token 缺失、无效或已过期 |
|
||||
|
||||
**conversation_id 校验失败**:
|
||||
|
||||
| 场景 | 处理 |
|
||||
|------|------|
|
||||
| 对话不存在 | 返回 401,`{"error": "SESSION_NOT_FOUND"}` |
|
||||
| 对话不属于当前用户 | 返回 401,`{"error": "SESSION_NOT_FOUND"}`(与不存在相同,避免信息泄露) |
|
||||
|
||||
**连接成功后**:`connected` 消息不变,新增 `conversation_id` 字段标识当前对话:
|
||||
|
||||
```typescript
|
||||
interface ConnectedMessage {
|
||||
type: "connected";
|
||||
session_id: string; // 对话 ID
|
||||
conversation_id: string; // 同 session_id,便于前端统一使用
|
||||
server_version: string;
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
### 四、前端调用示例
|
||||
|
||||
#### 认证状态管理
|
||||
|
||||
```typescript
|
||||
// 存储 token(建议 localStorage 或内存,视安全需求)
|
||||
interface AuthTokens {
|
||||
accessToken: string;
|
||||
refreshToken: string;
|
||||
}
|
||||
|
||||
// 请求拦截器:自动附加 Authorization 头
|
||||
async function authFetch(url: string, options: RequestInit = {}): Promise<Response> {
|
||||
const tokens = getStoredTokens();
|
||||
const headers = {
|
||||
...options.headers,
|
||||
"Authorization": `Bearer ${tokens.accessToken}`,
|
||||
};
|
||||
|
||||
let resp = await fetch(url, { ...options, headers });
|
||||
|
||||
// 401 时尝试刷新 token
|
||||
if (resp.status === 401 && tokens.refreshToken) {
|
||||
const refreshResp = await fetch("/api/auth/refresh", {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json" },
|
||||
body: JSON.stringify({ refresh_token: tokens.refreshToken }),
|
||||
});
|
||||
|
||||
if (refreshResp.ok) {
|
||||
const newTokens: AuthResponse = await refreshResp.json();
|
||||
storeTokens({
|
||||
accessToken: newTokens.access_token,
|
||||
refreshToken: newTokens.refresh_token,
|
||||
});
|
||||
// 用新 token 重试原请求
|
||||
headers["Authorization"] = `Bearer ${newTokens.access_token}`;
|
||||
resp = await fetch(url, { ...options, headers });
|
||||
} else {
|
||||
// refresh 也失败,跳转登录
|
||||
redirectToLogin();
|
||||
}
|
||||
}
|
||||
|
||||
return resp;
|
||||
}
|
||||
```
|
||||
|
||||
#### 注册 + 登录
|
||||
|
||||
```typescript
|
||||
async function register(username: string, password: string): Promise<AuthResponse> {
|
||||
const resp = await fetch("/api/auth/register", {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json" },
|
||||
body: JSON.stringify({ username, password }),
|
||||
});
|
||||
|
||||
if (!resp.ok) {
|
||||
const err: ApiError = await resp.json();
|
||||
throw new Error(err.message); // "username already taken" 等
|
||||
}
|
||||
|
||||
return resp.json();
|
||||
}
|
||||
```
|
||||
|
||||
#### 获取对话列表
|
||||
|
||||
```typescript
|
||||
async function getConversations(page = 1, size = 20): Promise<ConversationListResponse> {
|
||||
const resp = await authFetch(
|
||||
`/api/conversations?page=${page}&size=${size}`
|
||||
);
|
||||
if (!resp.ok) throw new Error("Failed to load conversations");
|
||||
return resp.json();
|
||||
}
|
||||
```
|
||||
|
||||
#### 加载对话历史消息
|
||||
|
||||
```typescript
|
||||
async function getMessages(
|
||||
conversationId: string,
|
||||
limit = 50,
|
||||
before?: number
|
||||
): Promise<MessagesResponse> {
|
||||
let url = `/api/conversations/${conversationId}/messages?limit=${limit}`;
|
||||
if (before !== undefined) url += `&before=${before}`;
|
||||
|
||||
const resp = await authFetch(url);
|
||||
if (!resp.ok) throw new Error("Failed to load messages");
|
||||
return resp.json();
|
||||
}
|
||||
```
|
||||
|
||||
#### 建立 WebSocket 连接(带认证)
|
||||
|
||||
```typescript
|
||||
function connectWebSocket(accessToken: string, conversationId?: string): WebSocket {
|
||||
let url = `/ws?token=${encodeURIComponent(accessToken)}`;
|
||||
if (conversationId) {
|
||||
url += `&conversation_id=${encodeURIComponent(conversationId)}`;
|
||||
}
|
||||
return new WebSocket(url);
|
||||
}
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 关键文件清单
|
||||
|
||||
```
|
||||
|
||||
Reference in New Issue
Block a user