Files
X-Agents/server/internal/handler/session_handler.go
DESKTOP-72TV0V4\caoxiaozhu 31f0feafb5 feat: 增强会话管理和 Agent 服务
- 优化 session_handler 会话处理逻辑
- 增强 agent_service Agent 服务功能
- 新增 chat_repository 仓储方法
- 更新 agent_handler 和 chat_group_handler
- 更新数据模型 agent 和 chat_session

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
2026-03-15 19:49:27 +08:00

328 lines
8.1 KiB
Go
Raw Permalink 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 handler
import (
"fmt"
"net/http"
"strconv"
"strings"
"x-agents/server/internal/model"
"x-agents/server/internal/repository"
"x-agents/server/internal/service"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
)
// ChatHandler 处理聊天请求
type ChatHandler struct {
chatService *service.ChatService
}
func NewChatHandler(chatService *service.ChatService) *ChatHandler {
return &ChatHandler{chatService: chatService}
}
// Chat 处理聊天请求
func (h *ChatHandler) Chat(c *gin.Context) {
var req model.AgentRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
// 从上下文获取用户ID由中间件设置
userID, exists := c.Get("user_id")
if !exists {
c.JSON(http.StatusUnauthorized, gin.H{"error": "unauthorized"})
return
}
resp, err := h.chatService.Chat(c.Request.Context(), userID.(string), req)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, resp)
}
// ListAgents 获取 Agent 列表
func (h *ChatHandler) ListAgents(c *gin.Context) {
userID, exists := c.Get("user_id")
if !exists {
c.JSON(http.StatusUnauthorized, gin.H{"error": "unauthorized"})
return
}
agents, err := h.chatService.ListAgents(userID.(string))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
if agents == nil {
agents = []model.Agent{}
}
c.JSON(http.StatusOK, gin.H{"agents": agents})
}
// CreateAgent 创建 Agent
func (h *ChatHandler) CreateAgent(c *gin.Context) {
userID, exists := c.Get("user_id")
if !exists {
c.JSON(http.StatusUnauthorized, gin.H{"error": "unauthorized"})
return
}
var req struct {
Name string `json:"name" binding:"required"`
Description string `json:"description"`
}
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
agent, err := h.chatService.CreateAgent(userID.(string), req.Name, req.Description)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusCreated, agent)
}
// SessionHandler 处理会话管理
type SessionHandler struct {
chatRepo *repository.ChatRepository
agentService *service.AgentService
}
func NewSessionHandler(chatRepo *repository.ChatRepository, agentService *service.AgentService) *SessionHandler {
return &SessionHandler{chatRepo: chatRepo, agentService: agentService}
}
// CreateSession 创建会话
func (h *SessionHandler) CreateSession(c *gin.Context) {
var req model.CreateSessionRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
session := &model.ChatSession{
ID: uuid.New().String(),
UserID: req.UserID,
AgentID: req.AgentID,
Title: req.Title,
ModelID: req.ModelID,
Status: "active",
}
if err := h.chatRepo.CreateSession(session); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, session)
}
// ListSessions 获取会话列表
func (h *SessionHandler) ListSessions(c *gin.Context) {
userID := c.Query("user_id")
if userID == "" {
c.JSON(http.StatusBadRequest, gin.H{"error": "user_id is required"})
return
}
limit, _ := strconv.Atoi(c.DefaultQuery("limit", "20"))
offset, _ := strconv.Atoi(c.DefaultQuery("offset", "0"))
sessions, total, err := h.chatRepo.GetSessionsByUserID(userID, limit, offset)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"list": sessions, "total": total})
}
// GetSession 获取会话详情
func (h *SessionHandler) GetSession(c *gin.Context) {
id := c.Param("id")
session, err := h.chatRepo.GetSessionByID(id)
if err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "Session not found"})
return
}
c.JSON(http.StatusOK, session)
}
// UpdateSession 更新会话
func (h *SessionHandler) UpdateSession(c *gin.Context) {
id := c.Param("id")
session, err := h.chatRepo.GetSessionByID(id)
if err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "Session not found"})
return
}
var req model.UpdateSessionRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if req.Title != "" {
session.Title = req.Title
}
if req.Status != "" {
session.Status = req.Status
}
if err := h.chatRepo.UpdateSession(session); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, session)
}
// DeleteSession 删除会话
func (h *SessionHandler) DeleteSession(c *gin.Context) {
id := c.Param("id")
if err := h.chatRepo.DeleteSession(id); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"success": true})
}
// GetMessages 获取会话消息
func (h *SessionHandler) GetMessages(c *gin.Context) {
sessionID := c.Param("id")
limit, _ := strconv.Atoi(c.DefaultQuery("limit", "100"))
offset, _ := strconv.Atoi(c.DefaultQuery("offset", "0"))
messages, total, err := h.chatRepo.GetMessagesBySessionID(sessionID, limit, offset)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"list": messages, "total": total})
}
// CreateMessage 创建消息
func (h *SessionHandler) CreateMessage(c *gin.Context) {
var req model.CreateMessageRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
// Debug: 打印请求内容
fmt.Printf("[CreateMessage] Request: session_id=%s, role=%s, content_len=%d\n",
req.SessionID, req.Role, len(req.Content))
// 验证 content 不为空
if req.Content == "" {
c.JSON(http.StatusBadRequest, gin.H{"error": "Content cannot be empty"})
return
}
// 检查会话是否存在
_, err := h.chatRepo.GetSessionByID(req.SessionID)
if err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "Session not found"})
return
}
message := &model.ChatMessage{
ID: uuid.New().String(),
SessionID: req.SessionID,
Role: req.Role,
Content: req.Content,
TokensUsed: req.TokensUsed,
DurationMs: req.DurationMs,
Metadata: req.Metadata,
}
if err := h.chatRepo.CreateMessage(message); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, message)
}
// GenerateSessionTitleRequest 生成会话标题请求
type GenerateSessionTitleRequest struct {
SessionID string `json:"session_id" binding:"required"`
}
// GenerateSessionTitle 生成会话标题
func (h *SessionHandler) GenerateSessionTitle(c *gin.Context) {
var req GenerateSessionTitleRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
// 获取会话的所有消息
messages, _, err := h.chatRepo.GetMessagesBySessionID(req.SessionID, 100, 0)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "Failed to get messages"})
return
}
if len(messages) < 2 {
c.JSON(http.StatusBadRequest, gin.H{"error": "Not enough messages to generate title"})
return
}
// 获取会话信息
session, err := h.chatRepo.GetSessionByID(req.SessionID)
if err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "Session not found"})
return
}
// 简单方式从用户的第一条消息中提取前15个字符作为标题
var title string
for _, msg := range messages {
if msg.Role == "user" {
// 清理内容,去除换行和多余空格
content := strings.ReplaceAll(msg.Content, "\n", " ")
content = strings.TrimSpace(content)
// 限制长度
if len(content) > 15 {
title = content[:15] + "..."
} else {
title = content
}
break
}
}
if title == "" {
title = "新会话"
}
fmt.Printf("[GenerateSessionTitle] Generated title: %s\n", title)
// 更新会话标题
session.Title = title
h.chatRepo.UpdateSession(session)
c.JSON(http.StatusOK, gin.H{"title": title})
}