feat: 实现 ConversationHandler(对话 CRUD + 消息查询 + 权限校验 + 路由注册)
- List: GET /api/conversations 获取当前用户对话列表(分页) - Create: POST /api/conversations 创建新对话 - Get: GET /api/conversations/:id 获取对话详情 - UpdateTitle: PATCH /api/conversations/:id 更新对话标题 - Delete: DELETE /api/conversations/:id 删除对话 - GetMessages: GET /api/conversations/:id/messages 获取消息列表 - 所有端点通过 AuthMiddleware 认证 - getSessionForUser 校验 session.UserID == claims.UserID - 返回 404 而非 403 避免信息泄露
This commit is contained in:
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
|
||||
}
|
||||
Reference in New Issue
Block a user