diff --git a/backend/internal/api/conversation.go b/backend/internal/api/conversation.go new file mode 100644 index 0000000..f95e2c1 --- /dev/null +++ b/backend/internal/api/conversation.go @@ -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 +}