Files
CamTalk/backend/internal/api/conversation.go
hhs 4cc0713459 fix: 修复空会话返回 null 导致前端崩溃的问题
## 问题描述
当用户创建新会话但未发送任何消息就切换到其他会话时,前端控制台报错:
"TypeError: Cannot read properties of null (reading 'map')"

根本原因:Go 后端未初始化的切片序列化为 JSON 时会变成 `null` 而非 `[]`,
前端尝试对 `null` 调用 `.map()` 导致崩溃。

## 修复方案
采用多层防御策略,同时修复后端和前端:

### 后端修复(确保 API 契约正确)
1. message_pg.go:93 - 将 `var messages []StoredMessage` 改为
   `messages := make([]StoredMessage, 0)`,确保空结果序列化为 `[]`
2. conversation.go - 在两个响应路径(PG 查询 + 内存回退)添加防御性 nil 检查

### 前端防御(多层保护)
1. useSessionList.ts - 在 loadMessages 和 loadSessions 中添加 null 合并操作
   `(res.data.messages || [])` 确保即使后端退化也不会崩溃

## 影响范围
- 所有空会话(新建后未发送消息的对话)现在可以正常切换
- API 响应符合 JSON 最佳实践(数组字段永远是 `[]` 而非 `null`)
2026-06-22 12:39:37 +08:00

359 lines
9.4 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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"
"github.com/hhs/camtalk/internal/store"
"github.com/hhs/camtalk/internal/trace"
)
// ConversationHandler 提供对话相关的 REST 端点。
type ConversationHandler struct {
sessionMgr session.Manager
tokenMgr *auth.TokenManager
msgRepo store.MessageRepository // 可选,为 nil 时 fallback 到内存查询
}
// NewConversationHandler 创建 ConversationHandler。
// msgRepo 可选,为 nil 时消息查询走内存。
func NewConversationHandler(sessionMgr session.Manager, tokenMgr *auth.TokenManager, msgRepo store.MessageRepository) *ConversationHandler {
return &ConversationHandler{
sessionMgr: sessionMgr,
tokenMgr: tokenMgr,
msgRepo: msgRepo,
}
}
// 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) {
log := trace.FromContext(c.Request.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 {
log.Errorw("list conversations failed",
"user_id", userID,
"error", err)
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) {
log := trace.FromContext(c.Request.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 {
log.Errorw("create conversation failed",
"user_id", userID,
"error", err)
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 {
log.Errorw("retrieve created conversation failed",
"session_id", sessionID,
"error", err)
c.JSON(http.StatusInternalServerError, gin.H{
"code": apperr.CodeInternalError,
"message": "failed to retrieve created conversation",
})
return
}
log.Infow("conversation created",
"conversation_id", sess.ID,
"user_id", userID)
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) {
log := trace.FromContext(c.Request.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
}
log.Errorw("update title failed",
"session_id", sessionID,
"error", err)
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) {
log := trace.FromContext(c.Request.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
}
log.Errorw("delete conversation failed",
"session_id", sessionID,
"error", err)
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: 消息 ID 游标(用于分页),返回此 ID 之前的消息
func (h *ConversationHandler) GetMessages(c *gin.Context) {
log := trace.FromContext(c.Request.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
}
beforeID, _ := strconv.ParseInt(c.DefaultQuery("before", "0"), 10, 64)
// 优先从 PostgreSQL 查询(支持持久化后的全量历史)
if h.msgRepo != nil {
messages, err := h.msgRepo.GetMessages(c.Request.Context(), sessionID, limit, beforeID)
if err != nil {
log.Errorw("get messages failed",
"session_id", sessionID,
"error", err)
c.JSON(http.StatusInternalServerError, gin.H{
"code": apperr.CodeInternalError,
"message": "failed to get messages",
})
return
}
count, _ := h.msgRepo.GetMessageCount(c.Request.Context(), sessionID)
if messages == nil {
messages = []store.StoredMessage{}
}
c.JSON(http.StatusOK, gin.H{
"messages": messages,
"total": count,
})
return
}
// fallback从内存查询
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
}
log.Errorw("get messages failed",
"session_id", sessionID,
"error", err)
c.JSON(http.StatusInternalServerError, gin.H{
"code": apperr.CodeInternalError,
"message": "failed to get messages",
})
return
}
total := len(allMessages)
// beforeID > 0 时表示偏移量(兼容旧接口语义)
if beforeID > 0 && int(beforeID) <= total {
allMessages = allMessages[:beforeID]
}
// 取最后 limit 条
start := len(allMessages) - limit
if start < 0 {
start = 0
}
messages := allMessages[start:]
if messages == nil {
messages = []models.Message{}
}
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
}