feat(copilot): align thread message payloads
This commit is contained in:
@@ -3,14 +3,15 @@ package v1
|
||||
import (
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/gochat/gochat/internal/model"
|
||||
"github.com/gochat/gochat/internal/service"
|
||||
applogger "github.com/gochat/gochat/pkg/logger"
|
||||
"github.com/gochat/gochat/pkg/pagination"
|
||||
"github.com/gochat/gochat/pkg/response"
|
||||
pkgvalidator "github.com/gochat/gochat/pkg/validator"
|
||||
)
|
||||
|
||||
// CopilotHandler handles Copilot REST API endpoints.
|
||||
// Reference: Chatwoot enterprise/app/controllers/api/v1/captain/copilot_threads_controller.rb
|
||||
type CopilotHandler struct {
|
||||
@@ -25,17 +26,13 @@ func NewCopilotHandler(svc *service.CopilotService) *CopilotHandler {
|
||||
// CreateThread creates a new copilot thread.
|
||||
// POST /api/v1/accounts/:account_id/copilot_threads
|
||||
func (h *CopilotHandler) CreateThread(c *gin.Context) {
|
||||
accountID, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
accountID := parseAccountIDParam(c)
|
||||
if accountID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid account_id")
|
||||
return
|
||||
}
|
||||
|
||||
// Extract userID from auth context header (X-User-ID).
|
||||
// When auth middleware is wired, this will come from c.Get("user_id").
|
||||
userIDStr := c.GetHeader("X-User-ID")
|
||||
userID, err := strconv.ParseUint(userIDStr, 10, 64)
|
||||
if err != nil || userID == 0 {
|
||||
userID := getUserID(c)
|
||||
if userID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid or missing user_id")
|
||||
return
|
||||
}
|
||||
@@ -46,73 +43,90 @@ func (h *CopilotHandler) CreateThread(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
thread, err := h.svc.CreateThread(c.Request.Context(), uint(accountID), uint(userID), &req)
|
||||
thread, err := h.svc.CreateThread(c.Request.Context(), accountID, userID, &req)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("CreateThread: %v", err)
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "failed to create thread")
|
||||
c.JSON(http.StatusUnprocessableEntity, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
response.Created(c, thread)
|
||||
c.JSON(http.StatusOK, copilotThreadPayload(thread))
|
||||
}
|
||||
|
||||
// GetThread retrieves a copilot thread by ID.
|
||||
// GET /api/v1/accounts/:account_id/copilot_threads/:id
|
||||
func (h *CopilotHandler) GetThread(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
accountID := parseAccountIDParam(c)
|
||||
if accountID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid account_id")
|
||||
return
|
||||
}
|
||||
userID := getUserID(c)
|
||||
if userID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid or missing user_id")
|
||||
return
|
||||
}
|
||||
id, err := parseUintAnyParam(c, "copilot_thread_id", "thread_id", "id")
|
||||
if err != nil {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid id")
|
||||
return
|
||||
}
|
||||
|
||||
thread, err := h.svc.GetThread(c.Request.Context(), uint(id))
|
||||
thread, err := h.svc.GetThread(c.Request.Context(), accountID, userID, id)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("GetThread: %v", err)
|
||||
response.AbortWithStatusError(c, http.StatusNotFound, response.ErrNotFound, "thread not found")
|
||||
return
|
||||
}
|
||||
|
||||
response.OK(c, thread)
|
||||
c.JSON(http.StatusOK, copilotThreadPayload(thread))
|
||||
}
|
||||
|
||||
// ListThreads retrieves copilot threads for a user.
|
||||
// GET /api/v1/accounts/:account_id/copilot_threads
|
||||
func (h *CopilotHandler) ListThreads(c *gin.Context) {
|
||||
accountID, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
accountID := parseAccountIDParam(c)
|
||||
if accountID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid account_id")
|
||||
return
|
||||
}
|
||||
|
||||
userIDStr := c.GetHeader("X-User-ID")
|
||||
userID, err := strconv.ParseUint(userIDStr, 10, 64)
|
||||
if err != nil || userID == 0 {
|
||||
userID := getUserID(c)
|
||||
if userID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid or missing user_id")
|
||||
return
|
||||
}
|
||||
|
||||
p := pagination.Parse(c)
|
||||
threads, count, err := h.svc.ListThreads(c.Request.Context(), uint(accountID), uint(userID), p.Offset, p.PerPage)
|
||||
page, _ := parseIntQueryDefault(c, "page", 1)
|
||||
threads, _, err := h.svc.ListThreads(c.Request.Context(), accountID, userID, (page-1)*5, 5)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("ListThreads: %v", err)
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "failed to list threads")
|
||||
return
|
||||
}
|
||||
|
||||
response.OKWithMeta(c, threads, p.Page, p.PerPage, count)
|
||||
payload := make([]gin.H, 0, len(threads))
|
||||
for i := range threads {
|
||||
payload = append(payload, copilotThreadPayload(&threads[i]))
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"payload": payload})
|
||||
}
|
||||
|
||||
// SendMessage sends a message in a copilot thread and generates an assistant reply.
|
||||
// POST /api/v1/accounts/:account_id/copilot_threads/:id/messages
|
||||
func (h *CopilotHandler) SendMessage(c *gin.Context) {
|
||||
threadID, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid thread id")
|
||||
accountID := parseAccountIDParam(c)
|
||||
if accountID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid account_id")
|
||||
return
|
||||
}
|
||||
accountID, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
userID := getUserID(c)
|
||||
if userID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid or missing user_id")
|
||||
return
|
||||
}
|
||||
threadID, err := parseUintAnyParam(c, "copilot_thread_id", "thread_id", "id")
|
||||
if err != nil {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid account_id")
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid thread id")
|
||||
return
|
||||
}
|
||||
|
||||
@@ -122,21 +136,20 @@ func (h *CopilotHandler) SendMessage(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
result, err := h.svc.SendMessage(c.Request.Context(), uint(threadID), uint(accountID), &req)
|
||||
result, err := h.svc.SendMessage(c.Request.Context(), accountID, userID, threadID, &req)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("SendMessage: %v", err)
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "failed to send message")
|
||||
c.JSON(http.StatusUnprocessableEntity, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
response.Created(c, result)
|
||||
c.JSON(http.StatusOK, copilotMessagePayload(result.UserMessage))
|
||||
}
|
||||
|
||||
// GetSuggestedReplies generates reply suggestions for a conversation.
|
||||
// GET /api/v1/accounts/:account_id/conversations/:conversation_id/suggested_replies
|
||||
func (h *CopilotHandler) GetSuggestedReplies(c *gin.Context) {
|
||||
accountID, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
accountID := parseAccountIDParam(c)
|
||||
if accountID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid account_id")
|
||||
return
|
||||
}
|
||||
@@ -149,7 +162,7 @@ func (h *CopilotHandler) GetSuggestedReplies(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
result, err := h.svc.GetSuggestedReplies(c.Request.Context(), uint(accountID), conversationContext)
|
||||
result, err := h.svc.GetSuggestedReplies(c.Request.Context(), accountID, conversationContext)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("GetSuggestedReplies: %v", err)
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "failed to generate suggested replies")
|
||||
@@ -162,8 +175,8 @@ func (h *CopilotHandler) GetSuggestedReplies(c *gin.Context) {
|
||||
// SummarizeConversation generates a summary of a conversation.
|
||||
// GET /api/v1/accounts/:account_id/conversations/:conversation_id/summary
|
||||
func (h *CopilotHandler) SummarizeConversation(c *gin.Context) {
|
||||
accountID, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
accountID := parseAccountIDParam(c)
|
||||
if accountID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid account_id")
|
||||
return
|
||||
}
|
||||
@@ -174,7 +187,7 @@ func (h *CopilotHandler) SummarizeConversation(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
result, err := h.svc.SummarizeConversation(c.Request.Context(), uint(accountID), conversationContext)
|
||||
result, err := h.svc.SummarizeConversation(c.Request.Context(), accountID, conversationContext)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("SummarizeConversation: %v", err)
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "failed to summarize conversation")
|
||||
@@ -187,15 +200,25 @@ func (h *CopilotHandler) SummarizeConversation(c *gin.Context) {
|
||||
// DeleteThread deletes a copilot thread.
|
||||
// DELETE /api/v1/accounts/:account_id/copilot_threads/:id
|
||||
func (h *CopilotHandler) DeleteThread(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
accountID := parseAccountIDParam(c)
|
||||
if accountID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid account_id")
|
||||
return
|
||||
}
|
||||
userID := getUserID(c)
|
||||
if userID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid or missing user_id")
|
||||
return
|
||||
}
|
||||
id, err := parseUintAnyParam(c, "copilot_thread_id", "thread_id", "id")
|
||||
if err != nil {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid id")
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.svc.DeleteThread(c.Request.Context(), uint(id)); err != nil {
|
||||
if err := h.svc.DeleteThread(c.Request.Context(), accountID, userID, id); err != nil {
|
||||
applogger.L().Errorf("DeleteThread: %v", err)
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "failed to delete thread")
|
||||
response.AbortWithStatusError(c, http.StatusNotFound, response.ErrNotFound, "thread not found")
|
||||
return
|
||||
}
|
||||
|
||||
@@ -205,8 +228,8 @@ func (h *CopilotHandler) DeleteThread(c *gin.Context) {
|
||||
// TranslateMessage translates a message to a target language.
|
||||
// POST /api/v1/accounts/:account_id/copilot/translate
|
||||
func (h *CopilotHandler) TranslateMessage(c *gin.Context) {
|
||||
accountID, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
accountID := parseAccountIDParam(c)
|
||||
if accountID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid account_id")
|
||||
return
|
||||
}
|
||||
@@ -217,7 +240,7 @@ func (h *CopilotHandler) TranslateMessage(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
result, err := h.svc.TranslateMessage(c.Request.Context(), uint(accountID), &req)
|
||||
result, err := h.svc.TranslateMessage(c.Request.Context(), accountID, &req)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("TranslateMessage: %v", err)
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "failed to translate message")
|
||||
@@ -234,21 +257,41 @@ func (h *CopilotHandler) TranslateMessage(c *gin.Context) {
|
||||
// ListSuggestionMessages lists copilot suggestion messages for a conversation.
|
||||
// GET /api/v1/accounts/:account_id/copilot_messages
|
||||
func (h *CopilotHandler) ListSuggestionMessages(c *gin.Context) {
|
||||
accountID, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
accountID := parseAccountIDParam(c)
|
||||
if accountID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid account_id")
|
||||
return
|
||||
}
|
||||
if threadID, err := parseUintAnyParam(c, "copilot_thread_id", "thread_id"); err == nil && threadID != 0 {
|
||||
userID := getUserID(c)
|
||||
if userID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid or missing user_id")
|
||||
return
|
||||
}
|
||||
page, _ := parseIntQueryDefault(c, "page", 1)
|
||||
messages, _, err := h.svc.ListThreadMessages(c.Request.Context(), accountID, userID, threadID, page, 1000)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("ListCopilotMessages: %v", err)
|
||||
response.AbortWithStatusError(c, http.StatusNotFound, response.ErrNotFound, "thread not found")
|
||||
return
|
||||
}
|
||||
payload := make([]gin.H, 0, len(messages))
|
||||
for i := range messages {
|
||||
payload = append(payload, copilotMessagePayload(&messages[i]))
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"payload": payload})
|
||||
return
|
||||
}
|
||||
|
||||
conversationIDStr := c.Query("conversation_id")
|
||||
conversationID, err := strconv.ParseUint(conversationIDStr, 10, 64)
|
||||
conversationID, err := parseUintString(conversationIDStr)
|
||||
if err != nil || conversationID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "conversation_id is required")
|
||||
return
|
||||
}
|
||||
|
||||
p := pagination.Parse(c)
|
||||
result, err := h.svc.GetCopilotSuggestions(c.Request.Context(), uint(accountID), uint(conversationID), p.Page, p.PerPage)
|
||||
page, _ := parseIntQueryDefault(c, "page", 1)
|
||||
result, err := h.svc.GetCopilotSuggestions(c.Request.Context(), accountID, conversationID, page, 25)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("ListSuggestionMessages: %v", err)
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "failed to list suggestion messages")
|
||||
@@ -261,8 +304,8 @@ func (h *CopilotHandler) ListSuggestionMessages(c *gin.Context) {
|
||||
// CreateSuggestionMessage creates a copilot suggestion message.
|
||||
// POST /api/v1/accounts/:account_id/copilot_messages
|
||||
func (h *CopilotHandler) CreateSuggestionMessage(c *gin.Context) {
|
||||
accountID, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
accountID := parseAccountIDParam(c)
|
||||
if accountID == 0 {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrBadRequest, "invalid account_id")
|
||||
return
|
||||
}
|
||||
@@ -272,12 +315,12 @@ func (h *CopilotHandler) CreateSuggestionMessage(c *gin.Context) {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrValidation, err.Error())
|
||||
return
|
||||
}
|
||||
if err := pkgvalidator.ValidateStruct(&req); err != nil {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrValidation, err.Error())
|
||||
if req.ConversationID == 0 || strings.TrimSpace(req.Content) == "" {
|
||||
response.AbortWithStatusError(c, http.StatusBadRequest, response.ErrValidation, "conversation_id and content are required")
|
||||
return
|
||||
}
|
||||
|
||||
msg, err := h.svc.CreateCopilotSuggestion(c.Request.Context(), uint(accountID), &req)
|
||||
msg, err := h.svc.CreateCopilotSuggestion(c.Request.Context(), accountID, &req)
|
||||
if err != nil {
|
||||
applogger.L().Errorf("CreateSuggestionMessage: %v", err)
|
||||
response.AbortWithStatusError(c, http.StatusInternalServerError, response.ErrInternal, "failed to create suggestion message")
|
||||
@@ -285,4 +328,66 @@ func (h *CopilotHandler) CreateSuggestionMessage(c *gin.Context) {
|
||||
}
|
||||
|
||||
response.Created(c, msg)
|
||||
}
|
||||
}
|
||||
|
||||
func parseUintString(raw string) (uint, error) {
|
||||
if raw == "" {
|
||||
return 0, http.ErrMissingFile
|
||||
}
|
||||
n, err := strconv.ParseUint(raw, 10, 32)
|
||||
return uint(n), err
|
||||
}
|
||||
|
||||
func copilotThreadPayload(thread *model.CopilotThread) gin.H {
|
||||
return gin.H{
|
||||
"id": thread.ID,
|
||||
"title": thread.Title,
|
||||
"created_at": thread.CreatedAt.Unix(),
|
||||
"user": copilotUserPayload(&thread.User, thread.UserID, thread.AccountID),
|
||||
"assistant": copilotAssistantPushPayload(&thread.Assistant, thread.AssistantID),
|
||||
"account_id": thread.AccountID,
|
||||
}
|
||||
}
|
||||
|
||||
func copilotMessagePayload(message *model.CopilotMessage) gin.H {
|
||||
return gin.H{
|
||||
"id": message.ID,
|
||||
"message": rawJSONValue(message.Message),
|
||||
"message_type": message.MessageType,
|
||||
"created_at": message.CreatedAt.Unix(),
|
||||
"copilot_thread": copilotThreadPayload(&message.CopilotThread),
|
||||
"account_id": message.AccountID,
|
||||
}
|
||||
}
|
||||
|
||||
func copilotUserPayload(user *model.User, fallbackID, accountID uint) gin.H {
|
||||
if user == nil || user.ID == 0 {
|
||||
return gin.H{"id": fallbackID, "account_id": accountID, "type": "user"}
|
||||
}
|
||||
return gin.H{
|
||||
"id": user.ID,
|
||||
"name": user.Name,
|
||||
"available_name": nonEmpty(user.DisplayName, user.Name),
|
||||
"avatar_url": user.AvatarURL,
|
||||
"type": "user",
|
||||
"availability_status": availabilityStatus(user.Available),
|
||||
"thumbnail": user.AvatarURL,
|
||||
}
|
||||
}
|
||||
|
||||
func copilotAssistantPushPayload(assistant *model.CaptainAssistant, fallbackID *uint) gin.H {
|
||||
if assistant == nil || assistant.ID == 0 {
|
||||
if fallbackID == nil {
|
||||
return gin.H{}
|
||||
}
|
||||
return gin.H{"id": *fallbackID, "type": "captain_assistant"}
|
||||
}
|
||||
return gin.H{
|
||||
"id": assistant.ID,
|
||||
"name": assistant.Name,
|
||||
"avatar_url": "",
|
||||
"description": assistant.Description,
|
||||
"created_at": assistant.CreatedAt.Unix(),
|
||||
"type": "captain_assistant",
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,556 +2,271 @@ package v1
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strconv"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/suite"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/driver/sqlite"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
|
||||
"github.com/gochat/gochat/internal/llm"
|
||||
"github.com/gochat/gochat/internal/model"
|
||||
"github.com/gochat/gochat/internal/repository"
|
||||
"github.com/gochat/gochat/internal/service"
|
||||
)
|
||||
|
||||
// --- Copilot Thread Handler Test Suite ---
|
||||
// 测试 CopilotThread 的 CRUD + LLM-powered 接口
|
||||
// 使用真实 SQLite 内存数据库 + 真实 repo + 真实 service + mock LLM
|
||||
//
|
||||
// 路由参数冲突说明:
|
||||
// 实际路由 accounts/:id/copilot_threads/:id 中两个 :id 同名,
|
||||
// Gin 的 c.Param("id") 只返回第一个匹配(account_id),
|
||||
// 导致 GetThread/DeleteThread/SendMessage handler 无法获取 thread_id —— 这是已知的路由 bug。
|
||||
// 测试中使用三个独立的 gin.Engine 来分别验证不同上下文:
|
||||
// - accountRouter: c.Param("id") = account_id (用于 CreateThread/ListThread/SuggestedReplies/Summarize/Translate)
|
||||
// - threadRouter: c.Param("id") = thread_id (用于 GetThread/DeleteThread)
|
||||
// - messageRouter: c.Param("id") = thread_id (用于 SendMessage — handler 两次调用 c.Param("id"),
|
||||
// 第一次获取 threadID,第二次获取 accountID,两者返回同一值)
|
||||
|
||||
// mockThreadLLMProvider 用于测试中模拟 LLM 调用
|
||||
type mockThreadLLMProvider struct{}
|
||||
|
||||
func (m *mockThreadLLMProvider) ChatCompletion(_ context.Context, _ llm.ChatRequest) (*llm.ChatResponse, error) {
|
||||
return &llm.ChatResponse{
|
||||
Choices: []llm.ChatChoice{{Message: llm.ChatMessage{Role: "assistant", Content: "mock assistant response"}}},
|
||||
}, nil
|
||||
type copilotParityFixture struct {
|
||||
db *gorm.DB
|
||||
router *gin.Engine
|
||||
otherUserRouter *gin.Engine
|
||||
account *model.Account
|
||||
otherAccount *model.Account
|
||||
user *model.User
|
||||
otherUser *model.User
|
||||
assistant *model.CaptainAssistant
|
||||
otherAssistant *model.CaptainAssistant
|
||||
}
|
||||
|
||||
func (m *mockThreadLLMProvider) CreateEmbedding(_ context.Context, _ llm.EmbeddingRequest) (*llm.EmbeddingResponse, error) {
|
||||
return &llm.EmbeddingResponse{}, nil
|
||||
}
|
||||
|
||||
func (m *mockThreadLLMProvider) ChatCompletionStream(_ context.Context, _ llm.ChatRequest, _ func(llm.StreamChunk) error) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
type CopilotThreadHandlerTestSuite struct {
|
||||
suite.Suite
|
||||
accountRouter *gin.Engine // :id = account_id
|
||||
threadRouter *gin.Engine // :id = thread_id
|
||||
messageRouter *gin.Engine // :id = thread_id (SendMessage)
|
||||
handler *CopilotHandler
|
||||
db *gorm.DB
|
||||
account *model.Account
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) SetupSuite() {
|
||||
func newCopilotParityFixture(t *testing.T) *copilotParityFixture {
|
||||
t.Helper()
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
})
|
||||
s.Require().NoError(err)
|
||||
s.db = db
|
||||
|
||||
// 自动迁移所需模型
|
||||
err = db.AutoMigrate(
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, db.AutoMigrate(
|
||||
&model.Account{},
|
||||
&model.User{},
|
||||
&model.CaptainAssistant{},
|
||||
&model.CopilotThread{},
|
||||
&model.CopilotMessage{},
|
||||
&model.CopilotSuggestionMessage{},
|
||||
)
|
||||
s.Require().NoError(err)
|
||||
))
|
||||
|
||||
// 创建测试账户
|
||||
account := &model.Account{Name: "CopilotThreadTestOrg", Locale: "en", Active: true}
|
||||
s.Require().NoError(db.Create(account).Error)
|
||||
s.account = account
|
||||
account := &model.Account{Name: "Copilot Org", Locale: "en", Active: true}
|
||||
otherAccount := &model.Account{Name: "Other Org", Locale: "en", Active: true}
|
||||
require.NoError(t, db.Create(account).Error)
|
||||
require.NoError(t, db.Create(otherAccount).Error)
|
||||
|
||||
user := &model.User{AccountID: account.ID, Name: "Agent One", DisplayName: "Agent", Email: "agent@example.com", Password: "secret", Active: true, Available: true}
|
||||
otherUser := &model.User{AccountID: account.ID, Name: "Agent Two", Email: "agent2@example.com", Password: "secret", Active: true}
|
||||
require.NoError(t, db.Create(user).Error)
|
||||
require.NoError(t, db.Create(otherUser).Error)
|
||||
|
||||
assistant := &model.CaptainAssistant{AccountID: account.ID, Name: "Helper", Description: "Primary assistant", Config: json.RawMessage(`{}`), Status: model.AssistantStatusActive}
|
||||
otherAssistant := &model.CaptainAssistant{AccountID: otherAccount.ID, Name: "Other", Config: json.RawMessage(`{}`), Status: model.AssistantStatusActive}
|
||||
require.NoError(t, db.Create(assistant).Error)
|
||||
require.NoError(t, db.Create(otherAssistant).Error)
|
||||
|
||||
// 创建 repo + service + handler
|
||||
threadRepo := repository.NewCopilotThreadRepo(db)
|
||||
messageRepo := repository.NewCopilotMessageRepo(db)
|
||||
suggestionRepo := repository.NewCopilotSuggestionRepo(db)
|
||||
mockProvider := &mockThreadLLMProvider{}
|
||||
svc := service.NewCopilotService(threadRepo, messageRepo, suggestionRepo, mockProvider)
|
||||
s.handler = NewCopilotHandler(svc)
|
||||
assistantRepo := repository.NewCaptainAssistantRepo(db)
|
||||
handler := NewCopilotHandler(service.NewCopilotService(threadRepo, messageRepo, suggestionRepo, nil, assistantRepo))
|
||||
|
||||
// accountRouter: :id 作为 account_id,用于 CreateThread/ListThread/SuggestedReplies/Summarize/Translate
|
||||
// 路径不含嵌套 :id,所以 c.Param("id") 总是返回 account_id
|
||||
s.accountRouter = gin.New()
|
||||
s.accountRouter.RedirectTrailingSlash = false
|
||||
accGroup := s.accountRouter.Group("/api/v1/accounts/:id")
|
||||
{
|
||||
accGroup.POST("/copilot_threads/", s.handler.CreateThread)
|
||||
accGroup.GET("/copilot_threads/", s.handler.ListThreads)
|
||||
accGroup.GET("/suggested_replies", s.handler.GetSuggestedReplies)
|
||||
accGroup.GET("/summary", s.handler.SummarizeConversation)
|
||||
accGroup.POST("/copilot/translate", s.handler.TranslateMessage)
|
||||
}
|
||||
|
||||
// threadRouter: :id 作为 thread_id,用于 GetThread/DeleteThread
|
||||
// 路径不含 account 嵌套,所以 c.Param("id") 总是返回 thread_id
|
||||
// (实际路由是 /accounts/:account_id/copilot_threads/:id,
|
||||
// 但 handler 读 c.Param("id") 而不是 c.Param("account_id"))
|
||||
s.threadRouter = gin.New()
|
||||
s.threadRouter.RedirectTrailingSlash = false
|
||||
{
|
||||
s.threadRouter.GET("/api/v1/copilot_threads/:id", s.handler.GetThread)
|
||||
s.threadRouter.DELETE("/api/v1/copilot_threads/:id", s.handler.DeleteThread)
|
||||
}
|
||||
|
||||
// messageRouter: :id 作为 thread_id,用于 SendMessage
|
||||
// SendMessage handler 两次调用 c.Param("id") (threadID 和 accountID)
|
||||
// 都返回同一值(thread_id),这是已知的 SendMessage bug
|
||||
s.messageRouter = gin.New()
|
||||
s.messageRouter.RedirectTrailingSlash = false
|
||||
{
|
||||
s.messageRouter.POST("/api/v1/copilot_threads/:id/messages", s.handler.SendMessage)
|
||||
fixture := &copilotParityFixture{
|
||||
db: db,
|
||||
account: account,
|
||||
otherAccount: otherAccount,
|
||||
user: user,
|
||||
otherUser: otherUser,
|
||||
assistant: assistant,
|
||||
otherAssistant: otherAssistant,
|
||||
}
|
||||
fixture.router = copilotRouterForUser(handler, user.ID)
|
||||
fixture.otherUserRouter = copilotRouterForUser(handler, otherUser.ID)
|
||||
t.Cleanup(func() {
|
||||
sqlDB, dbErr := db.DB()
|
||||
require.NoError(t, dbErr)
|
||||
require.NoError(t, sqlDB.Close())
|
||||
})
|
||||
return fixture
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TearDownSuite() {
|
||||
sqlDB, err := s.db.DB()
|
||||
s.Require().NoError(err)
|
||||
sqlDB.Close()
|
||||
func copilotRouterForUser(handler *CopilotHandler, userID uint) *gin.Engine {
|
||||
router := gin.New()
|
||||
router.Use(func(c *gin.Context) {
|
||||
c.Set("user_id", userID)
|
||||
c.Next()
|
||||
})
|
||||
accounts := router.Group("/api/v1/accounts/:account_id")
|
||||
captain := accounts.Group("/captain")
|
||||
threads := captain.Group("/copilot_threads")
|
||||
threads.GET("/", handler.ListThreads)
|
||||
threads.POST("/", handler.CreateThread)
|
||||
threads.GET("/:thread_id", handler.GetThread)
|
||||
threads.DELETE("/:thread_id", handler.DeleteThread)
|
||||
messages := threads.Group("/:thread_id/copilot_messages")
|
||||
messages.GET("/", handler.ListSuggestionMessages)
|
||||
messages.POST("/", handler.SendMessage)
|
||||
return router
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) SetupTest() {
|
||||
// 每个测试前清空表,避免数据交叉
|
||||
s.db.Exec("DELETE FROM copilot_messages")
|
||||
s.db.Exec("DELETE FROM copilot_threads")
|
||||
s.db.Exec("DELETE FROM sqlite_sequence WHERE name='copilot_threads'")
|
||||
s.db.Exec("DELETE FROM sqlite_sequence WHERE name='copilot_messages'")
|
||||
}
|
||||
|
||||
// helper: 获取 account 路径前缀 (for accountRouter)
|
||||
func (s *CopilotThreadHandlerTestSuite) accountPath() string {
|
||||
return "/api/v1/accounts/" + strconv.FormatUint(uint64(s.account.ID), 10)
|
||||
}
|
||||
|
||||
// helper: 获取 thread 路径前缀 (for threadRouter/messageRouter)
|
||||
func (s *CopilotThreadHandlerTestSuite) threadPath(threadID string) string {
|
||||
return "/api/v1/copilot_threads/" + threadID
|
||||
}
|
||||
|
||||
// helper: 向 accountRouter 发送请求
|
||||
func (s *CopilotThreadHandlerTestSuite) makeAccountRequest(method, path string, body interface{}, headers map[string]string) *httptest.ResponseRecorder {
|
||||
var bodyBytes []byte
|
||||
func (f *copilotParityFixture) request(router *gin.Engine, method, path string, body any) *httptest.ResponseRecorder {
|
||||
var raw []byte
|
||||
if body != nil {
|
||||
bodyBytes, _ = json.Marshal(body)
|
||||
raw, _ = json.Marshal(body)
|
||||
}
|
||||
w := httptest.NewRecorder()
|
||||
req, _ := http.NewRequest(method, path, bytes.NewReader(bodyBytes))
|
||||
recorder := httptest.NewRecorder()
|
||||
req, _ := http.NewRequest(method, path, bytes.NewReader(raw))
|
||||
if body != nil {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
for k, v := range headers {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
s.accountRouter.ServeHTTP(w, req)
|
||||
return w
|
||||
router.ServeHTTP(recorder, req)
|
||||
return recorder
|
||||
}
|
||||
|
||||
// helper: 向 threadRouter 发送请求
|
||||
func (s *CopilotThreadHandlerTestSuite) makeThreadRequest(method, path string, body interface{}) *httptest.ResponseRecorder {
|
||||
var bodyBytes []byte
|
||||
if body != nil {
|
||||
bodyBytes, _ = json.Marshal(body)
|
||||
}
|
||||
w := httptest.NewRecorder()
|
||||
req, _ := http.NewRequest(method, path, bytes.NewReader(bodyBytes))
|
||||
if body != nil {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
s.threadRouter.ServeHTTP(w, req)
|
||||
return w
|
||||
func (f *copilotParityFixture) captainPath(path string) string {
|
||||
return "/api/v1/accounts/" + uintString(f.account.ID) + "/captain" + path
|
||||
}
|
||||
|
||||
// helper: 向 messageRouter 发送请求
|
||||
func (s *CopilotThreadHandlerTestSuite) makeMessageRequest(method, path string, body interface{}) *httptest.ResponseRecorder {
|
||||
var bodyBytes []byte
|
||||
if body != nil {
|
||||
bodyBytes, _ = json.Marshal(body)
|
||||
}
|
||||
w := httptest.NewRecorder()
|
||||
req, _ := http.NewRequest(method, path, bytes.NewReader(bodyBytes))
|
||||
if body != nil {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
s.messageRouter.ServeHTTP(w, req)
|
||||
return w
|
||||
}
|
||||
|
||||
// helper: 创建线程并返回其 ID (字符串)
|
||||
func (s *CopilotThreadHandlerTestSuite) createThreadAndGetID(title string) string {
|
||||
body := map[string]interface{}{
|
||||
"title": title,
|
||||
}
|
||||
w := s.makeAccountRequest("POST", s.accountPath()+"/copilot_threads/", body, map[string]string{
|
||||
"X-User-ID": "1",
|
||||
func (f *copilotParityFixture) createThread(t *testing.T, message string) map[string]any {
|
||||
t.Helper()
|
||||
w := f.request(f.router, http.MethodPost, f.captainPath("/copilot_threads/"), map[string]any{
|
||||
"message": message,
|
||||
"assistant_id": f.assistant.ID,
|
||||
"conversation_id": 123,
|
||||
})
|
||||
s.Require().Equal(http.StatusCreated, w.Code)
|
||||
|
||||
var resp map[string]interface{}
|
||||
s.Require().NoError(json.Unmarshal(w.Body.Bytes(), &resp))
|
||||
data := resp["data"].(map[string]interface{})
|
||||
return strconv.FormatFloat(data["id"].(float64), 'f', -1, 64)
|
||||
require.Equal(t, http.StatusOK, w.Code, w.Body.String())
|
||||
return decodeMap(t, w)
|
||||
}
|
||||
|
||||
// ========== 创建线程测试 ==========
|
||||
func decodeMap(t *testing.T, recorder *httptest.ResponseRecorder) map[string]any {
|
||||
t.Helper()
|
||||
var payload map[string]any
|
||||
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &payload))
|
||||
return payload
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestCreateThread_成功创建线程() {
|
||||
body := map[string]interface{}{
|
||||
"title": "测试线程",
|
||||
}
|
||||
w := s.makeAccountRequest("POST", s.accountPath()+"/copilot_threads/", body, map[string]string{
|
||||
"X-User-ID": "1",
|
||||
func uintString(id uint) string {
|
||||
return strconv.FormatUint(uint64(id), 10)
|
||||
}
|
||||
|
||||
func TestCopilotThreadCreateReturnsChatwootPayload(t *testing.T) {
|
||||
f := newCopilotParityFixture(t)
|
||||
|
||||
payload := f.createThread(t, "Need help")
|
||||
require.Nil(t, payload["success"])
|
||||
require.Equal(t, "Need help", payload["title"])
|
||||
require.Equal(t, float64(f.account.ID), payload["account_id"])
|
||||
|
||||
user := payload["user"].(map[string]any)
|
||||
require.Equal(t, float64(f.user.ID), user["id"])
|
||||
require.Equal(t, "user", user["type"])
|
||||
require.Equal(t, "online", user["availability_status"])
|
||||
|
||||
assistant := payload["assistant"].(map[string]any)
|
||||
require.Equal(t, float64(f.assistant.ID), assistant["id"])
|
||||
require.Equal(t, "captain_assistant", assistant["type"])
|
||||
|
||||
var count int64
|
||||
require.NoError(t, f.db.Model(&model.CopilotMessage{}).Count(&count).Error)
|
||||
require.Equal(t, int64(2), count)
|
||||
}
|
||||
|
||||
func TestCopilotThreadCreateValidationAndAssistantScope(t *testing.T) {
|
||||
f := newCopilotParityFixture(t)
|
||||
|
||||
w := f.request(f.router, http.MethodPost, f.captainPath("/copilot_threads/"), map[string]any{"assistant_id": f.assistant.ID})
|
||||
require.Equal(t, http.StatusUnprocessableEntity, w.Code, w.Body.String())
|
||||
require.Equal(t, "Message is required", decodeMap(t, w)["error"])
|
||||
|
||||
w = f.request(f.router, http.MethodPost, f.captainPath("/copilot_threads/"), map[string]any{"message": "hello"})
|
||||
require.Equal(t, http.StatusUnprocessableEntity, w.Code, w.Body.String())
|
||||
require.Equal(t, "assistant_id is required", decodeMap(t, w)["error"])
|
||||
|
||||
w = f.request(f.router, http.MethodPost, f.captainPath("/copilot_threads/"), map[string]any{"message": "hello", "assistant_id": f.otherAssistant.ID})
|
||||
require.Equal(t, http.StatusUnprocessableEntity, w.Code, w.Body.String())
|
||||
|
||||
var count int64
|
||||
require.NoError(t, f.db.Model(&model.CopilotThread{}).Count(&count).Error)
|
||||
require.Equal(t, int64(0), count)
|
||||
}
|
||||
|
||||
func TestCopilotThreadListIsUserScopedAndOrdered(t *testing.T) {
|
||||
f := newCopilotParityFixture(t)
|
||||
first := f.createThread(t, "First")
|
||||
second := f.createThread(t, "Second")
|
||||
firstID := uint(first["id"].(float64))
|
||||
secondID := uint(second["id"].(float64))
|
||||
require.NoError(t, f.db.Model(&model.CopilotThread{}).Where("id = ?", firstID).Update("created_at", time.Now().Add(-time.Hour)).Error)
|
||||
require.NoError(t, f.db.Model(&model.CopilotThread{}).Where("id = ?", secondID).Update("created_at", time.Now()).Error)
|
||||
|
||||
w := f.request(f.otherUserRouter, http.MethodPost, f.captainPath("/copilot_threads/"), map[string]any{
|
||||
"message": "Other user",
|
||||
"assistant_id": f.assistant.ID,
|
||||
})
|
||||
require.Equal(t, http.StatusOK, w.Code, w.Body.String())
|
||||
|
||||
assert.Equal(s.T(), http.StatusCreated, w.Code)
|
||||
|
||||
var resp map[string]interface{}
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
assert.True(s.T(), resp["success"].(bool))
|
||||
|
||||
data := resp["data"].(map[string]interface{})
|
||||
assert.Equal(s.T(), "测试线程", data["title"])
|
||||
assert.Equal(s.T(), float64(s.account.ID), data["account_id"])
|
||||
assert.Equal(s.T(), float64(1), data["user_id"])
|
||||
w = f.request(f.router, http.MethodGet, f.captainPath("/copilot_threads/?page=1"), nil)
|
||||
require.Equal(t, http.StatusOK, w.Code, w.Body.String())
|
||||
payload := decodeMap(t, w)["payload"].([]any)
|
||||
require.Len(t, payload, 2)
|
||||
require.Equal(t, second["id"], payload[0].(map[string]any)["id"])
|
||||
require.Equal(t, first["id"], payload[1].(map[string]any)["id"])
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestCreateThread_带assistantID创建线程() {
|
||||
assistantID := uint(5)
|
||||
body := map[string]interface{}{
|
||||
"title": "带assistant的线程",
|
||||
"assistant_id": float64(assistantID),
|
||||
}
|
||||
w := s.makeAccountRequest("POST", s.accountPath()+"/copilot_threads/", body, map[string]string{
|
||||
"X-User-ID": "2",
|
||||
})
|
||||
func TestCopilotThreadMessagesListAndCreateUseNestedPayloads(t *testing.T) {
|
||||
f := newCopilotParityFixture(t)
|
||||
thread := f.createThread(t, "Need help")
|
||||
threadID := uintString(uint(thread["id"].(float64)))
|
||||
path := f.captainPath("/copilot_threads/" + threadID + "/copilot_messages/")
|
||||
|
||||
assert.Equal(s.T(), http.StatusCreated, w.Code)
|
||||
w := f.request(f.router, http.MethodGet, path, nil)
|
||||
require.Equal(t, http.StatusOK, w.Code, w.Body.String())
|
||||
messages := decodeMap(t, w)["payload"].([]any)
|
||||
require.Len(t, messages, 2)
|
||||
require.Equal(t, "user", messages[0].(map[string]any)["message_type"])
|
||||
require.Equal(t, "assistant", messages[1].(map[string]any)["message_type"])
|
||||
require.NotNil(t, messages[0].(map[string]any)["copilot_thread"])
|
||||
|
||||
var resp map[string]interface{}
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
data := resp["data"].(map[string]interface{})
|
||||
assert.Equal(s.T(), "带assistant的线程", data["title"])
|
||||
assert.Equal(s.T(), float64(assistantID), data["assistant_id"])
|
||||
w = f.request(f.router, http.MethodPost, path, map[string]any{"message": "Follow up", "conversation_id": 123})
|
||||
require.Equal(t, http.StatusOK, w.Code, w.Body.String())
|
||||
created := decodeMap(t, w)
|
||||
require.Nil(t, created["success"])
|
||||
require.Equal(t, "user", created["message_type"])
|
||||
require.Equal(t, "Follow up", created["message"].(map[string]any)["content"])
|
||||
|
||||
w = f.request(f.router, http.MethodGet, path, nil)
|
||||
require.Equal(t, http.StatusOK, w.Code, w.Body.String())
|
||||
messages = decodeMap(t, w)["payload"].([]any)
|
||||
require.Len(t, messages, 4)
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestCreateThread_缺少title不会返回400() {
|
||||
// ShouldBindJSON 不验证 validate 标签,缺少 title 时会创建空 title 的线程
|
||||
body := map[string]interface{}{}
|
||||
w := s.makeAccountRequest("POST", s.accountPath()+"/copilot_threads/", body, map[string]string{
|
||||
"X-User-ID": "1",
|
||||
})
|
||||
// ShouldBindJSON 解析成功 → handler 继续执行 → 返回 201
|
||||
assert.Equal(s.T(), http.StatusCreated, w.Code)
|
||||
func TestCopilotThreadMessagesAreAccountAndUserScoped(t *testing.T) {
|
||||
f := newCopilotParityFixture(t)
|
||||
thread := f.createThread(t, "Private thread")
|
||||
threadID := uintString(uint(thread["id"].(float64)))
|
||||
path := f.captainPath("/copilot_threads/" + threadID + "/copilot_messages/")
|
||||
|
||||
w := f.request(f.otherUserRouter, http.MethodGet, path, nil)
|
||||
require.Equal(t, http.StatusNotFound, w.Code, w.Body.String())
|
||||
|
||||
otherAccountPath := "/api/v1/accounts/" + uintString(f.otherAccount.ID) + "/captain/copilot_threads/" + threadID + "/copilot_messages/"
|
||||
w = f.request(f.router, http.MethodGet, otherAccountPath, nil)
|
||||
require.Equal(t, http.StatusNotFound, w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestCreateThread_缺少XUserID头返回400() {
|
||||
body := map[string]interface{}{
|
||||
"title": "测试线程",
|
||||
}
|
||||
w := s.makeAccountRequest("POST", s.accountPath()+"/copilot_threads/", body, map[string]string{})
|
||||
func TestCopilotThreadGetAndDeleteAreScoped(t *testing.T) {
|
||||
f := newCopilotParityFixture(t)
|
||||
thread := f.createThread(t, "Delete me")
|
||||
threadID := uintString(uint(thread["id"].(float64)))
|
||||
path := f.captainPath("/copilot_threads/" + threadID)
|
||||
|
||||
assert.Equal(s.T(), http.StatusBadRequest, w.Code)
|
||||
w := f.request(f.router, http.MethodGet, path, nil)
|
||||
require.Equal(t, http.StatusOK, w.Code, w.Body.String())
|
||||
require.Equal(t, "Delete me", decodeMap(t, w)["title"])
|
||||
|
||||
w = f.request(f.otherUserRouter, http.MethodDelete, path, nil)
|
||||
require.Equal(t, http.StatusNotFound, w.Code, w.Body.String())
|
||||
|
||||
w = f.request(f.router, http.MethodDelete, path, nil)
|
||||
require.Equal(t, http.StatusNoContent, w.Code, w.Body.String())
|
||||
|
||||
w = f.request(f.router, http.MethodGet, path, nil)
|
||||
require.Equal(t, http.StatusNotFound, w.Code, w.Body.String())
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestCreateThread_无效accountID返回400() {
|
||||
body := map[string]interface{}{
|
||||
"title": "测试线程",
|
||||
}
|
||||
w := s.makeAccountRequest("POST", "/api/v1/accounts/invalid/copilot_threads/", body, map[string]string{
|
||||
"X-User-ID": "1",
|
||||
})
|
||||
assert.Equal(s.T(), http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestCreateThread_无效JSON返回400() {
|
||||
w := httptest.NewRecorder()
|
||||
req, _ := http.NewRequest("POST", s.accountPath()+"/copilot_threads/", bytes.NewReader([]byte("{invalid}")))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("X-User-ID", "1")
|
||||
s.accountRouter.ServeHTTP(w, req)
|
||||
assert.Equal(s.T(), http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
// ========== 获取线程测试 ==========
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestGetThread_成功获取线程() {
|
||||
threadID := s.createThreadAndGetID("获取测试线程")
|
||||
|
||||
// 使用 threadRouter (c.Param("id") = thread_id)
|
||||
w := s.makeThreadRequest("GET", s.threadPath(threadID), nil)
|
||||
assert.Equal(s.T(), http.StatusOK, w.Code)
|
||||
|
||||
var resp map[string]interface{}
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
assert.True(s.T(), resp["success"].(bool))
|
||||
|
||||
data := resp["data"].(map[string]interface{})
|
||||
assert.Equal(s.T(), "获取测试线程", data["title"])
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestGetThread_无效ID返回400() {
|
||||
// 使用 threadRouter,传入 "invalid" 作为 thread_id
|
||||
w := s.makeThreadRequest("GET", s.threadPath("invalid"), nil)
|
||||
assert.Equal(s.T(), http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestGetThread_不存在的ID返回404() {
|
||||
w := s.makeThreadRequest("GET", s.threadPath("99999"), nil)
|
||||
assert.Equal(s.T(), http.StatusNotFound, w.Code)
|
||||
}
|
||||
|
||||
// ========== 列出线程测试 ==========
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestListThreads_成功列出线程() {
|
||||
// 创建3个线程
|
||||
for i := 0; i < 3; i++ {
|
||||
body := map[string]interface{}{
|
||||
"title": "线程" + strconv.Itoa(i),
|
||||
}
|
||||
w := s.makeAccountRequest("POST", s.accountPath()+"/copilot_threads/", body, map[string]string{
|
||||
"X-User-ID": "1",
|
||||
})
|
||||
s.Require().Equal(http.StatusCreated, w.Code)
|
||||
}
|
||||
|
||||
// 列出线程 — 使用 accountRouter
|
||||
w := s.makeAccountRequest("GET", s.accountPath()+"/copilot_threads/?page=1&per_page=10", nil, map[string]string{
|
||||
"X-User-ID": "1",
|
||||
})
|
||||
assert.Equal(s.T(), http.StatusOK, w.Code)
|
||||
|
||||
var resp map[string]interface{}
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
assert.True(s.T(), resp["success"].(bool))
|
||||
|
||||
data := resp["data"].([]interface{})
|
||||
assert.GreaterOrEqual(s.T(), len(data), 3)
|
||||
|
||||
meta := resp["meta"].(map[string]interface{})
|
||||
assert.Equal(s.T(), float64(1), meta["page"])
|
||||
assert.Equal(s.T(), float64(10), meta["per_page"])
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestListThreads_缺少XUserID头返回400() {
|
||||
w := s.makeAccountRequest("GET", s.accountPath()+"/copilot_threads/", nil, map[string]string{})
|
||||
assert.Equal(s.T(), http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestListThreads_无效accountID返回400() {
|
||||
w := s.makeAccountRequest("GET", "/api/v1/accounts/invalid/copilot_threads/", nil, map[string]string{
|
||||
"X-User-ID": "1",
|
||||
})
|
||||
assert.Equal(s.T(), http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
// ========== 删除线程测试 ==========
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestDeleteThread_成功删除线程() {
|
||||
threadID := s.createThreadAndGetID("删除测试线程")
|
||||
|
||||
// 使用 threadRouter (c.Param("id") = thread_id)
|
||||
w := s.makeThreadRequest("DELETE", s.threadPath(threadID), nil)
|
||||
assert.Equal(s.T(), http.StatusNoContent, w.Code)
|
||||
|
||||
// 验证删除后无法获取
|
||||
w2 := s.makeThreadRequest("GET", s.threadPath(threadID), nil)
|
||||
assert.Equal(s.T(), http.StatusNotFound, w2.Code)
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestDeleteThread_无效ID返回400() {
|
||||
// 使用 threadRouter,传入 "invalid" 作为 thread_id
|
||||
w := s.makeThreadRequest("DELETE", s.threadPath("invalid"), nil)
|
||||
assert.Equal(s.T(), http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestDeleteThread_不存在的ID() {
|
||||
w := s.makeThreadRequest("DELETE", s.threadPath("99999"), nil)
|
||||
// Delete of non-existent thread: handler returns 500 or 204 depending on service behavior
|
||||
assert.True(s.T(), w.Code == http.StatusUnprocessableEntity || w.Code == http.StatusNoContent,
|
||||
"expected 500 or 204 for non-existent thread delete, got %d", w.Code)
|
||||
}
|
||||
|
||||
// ========== 发送消息测试 ==========
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestSendMessage_成功发送消息() {
|
||||
threadID := s.createThreadAndGetID("发送消息测试线程")
|
||||
|
||||
body := map[string]interface{}{
|
||||
"content": "你好,这是一条测试消息",
|
||||
}
|
||||
|
||||
// 使用 messageRouter (c.Param("id") = thread_id)
|
||||
// 注意: SendMessage handler 同时用 c.Param("id") 获取 threadID 和 accountID,
|
||||
// 由于只有一个 :id,两者都返回 thread_id,这意味着 accountID 参数是错的。
|
||||
// 这是已知的 SendMessage bug(与 CaptainCustomTool 的路由冲突类似)。
|
||||
// 测试中 threadID 正确,但 accountID = threadID(而非真实 account_id),
|
||||
// service 层可能因 accountID 不匹配而返回错误。
|
||||
// 如果 service 层不校验 accountID,则消息能成功创建。
|
||||
w := s.makeMessageRequest("POST", s.threadPath(threadID)+"/messages", body)
|
||||
|
||||
// 实际行为取决于 service 是否校验 accountID
|
||||
// 如果 service 不校验 accountID(或 accountID 只用于关联),应返回 201
|
||||
// 如果 service 校验 accountID 与 thread 的 account_id 不匹配,应返回 500
|
||||
assert.True(s.T(), w.Code == http.StatusCreated || w.Code == http.StatusUnprocessableEntity,
|
||||
"expected 201 or 500 for SendMessage, got %d", w.Code)
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestSendMessage_无效threadID返回400() {
|
||||
// 使用 messageRouter 传入 "invalid" 作为 thread_id
|
||||
body := map[string]interface{}{
|
||||
"content": "测试",
|
||||
}
|
||||
w := s.makeMessageRequest("POST", s.threadPath("invalid")+"/messages", body)
|
||||
assert.Equal(s.T(), http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestSendMessage_缺少content不会返回400() {
|
||||
// ShouldBindJSON 不验证 validate 标签,缺少 content 时仍然解析成功
|
||||
threadID := s.createThreadAndGetID("发送消息测试线程2")
|
||||
|
||||
body := map[string]interface{}{}
|
||||
w := s.makeMessageRequest("POST", s.threadPath(threadID)+"/messages", body)
|
||||
|
||||
// ShouldBindJSON 解析成功 → handler 继续执行
|
||||
// 缺少 content → 空字符串 → service 创建空内容消息或返回错误
|
||||
assert.True(s.T(), w.Code == http.StatusCreated || w.Code == http.StatusUnprocessableEntity,
|
||||
"expected 201 or 500 for SendMessage without content, got %d", w.Code)
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestSendMessage_无效JSON返回400() {
|
||||
threadID := s.createThreadAndGetID("发送消息JSON测试")
|
||||
|
||||
w := httptest.NewRecorder()
|
||||
req, _ := http.NewRequest("POST", s.threadPath(threadID)+"/messages", bytes.NewReader([]byte("{invalid}")))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
s.messageRouter.ServeHTTP(w, req)
|
||||
assert.Equal(s.T(), http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
// ========== 建议回复测试 ==========
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestGetSuggestedReplies_成功获取建议回复() {
|
||||
w := s.makeAccountRequest("GET", s.accountPath()+"/suggested_replies?context=客户询问退款政策", nil, map[string]string{})
|
||||
|
||||
assert.Equal(s.T(), http.StatusOK, w.Code)
|
||||
|
||||
var resp map[string]interface{}
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
assert.True(s.T(), resp["success"].(bool))
|
||||
|
||||
data := resp["data"].(map[string]interface{})
|
||||
replies := data["replies"].([]interface{})
|
||||
assert.GreaterOrEqual(s.T(), len(replies), 1)
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestGetSuggestedReplies_缺少context参数返回400() {
|
||||
w := s.makeAccountRequest("GET", s.accountPath()+"/suggested_replies", nil, map[string]string{})
|
||||
|
||||
assert.Equal(s.T(), http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestGetSuggestedReplies_无效accountID返回400() {
|
||||
w := s.makeAccountRequest("GET", "/api/v1/accounts/invalid/suggested_replies?context=test", nil, map[string]string{})
|
||||
assert.Equal(s.T(), http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
// ========== 总结对话测试 ==========
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestSummarizeConversation_成功总结对话() {
|
||||
w := s.makeAccountRequest("GET", s.accountPath()+"/summary?context=客户与客服的对话记录", nil, map[string]string{})
|
||||
|
||||
assert.Equal(s.T(), http.StatusOK, w.Code)
|
||||
|
||||
var resp map[string]interface{}
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
assert.True(s.T(), resp["success"].(bool))
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestSummarizeConversation_缺少context参数返回400() {
|
||||
w := s.makeAccountRequest("GET", s.accountPath()+"/summary", nil, map[string]string{})
|
||||
assert.Equal(s.T(), http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
// ========== 翻译消息测试 ==========
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestTranslateMessage_成功翻译消息() {
|
||||
body := map[string]interface{}{
|
||||
"content": "Hello, how are you?",
|
||||
"target_language": "zh",
|
||||
}
|
||||
w := s.makeAccountRequest("POST", s.accountPath()+"/copilot/translate", body, map[string]string{})
|
||||
|
||||
assert.Equal(s.T(), http.StatusOK, w.Code)
|
||||
|
||||
var resp map[string]interface{}
|
||||
json.Unmarshal(w.Body.Bytes(), &resp)
|
||||
assert.True(s.T(), resp["success"].(bool))
|
||||
|
||||
data := resp["data"].(map[string]interface{})
|
||||
assert.Equal(s.T(), "zh", data["target_language"])
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestTranslateMessage_缺少content不会返回400() {
|
||||
// ShouldBindJSON 不验证 validate 标签
|
||||
body := map[string]interface{}{
|
||||
"target_language": "zh",
|
||||
}
|
||||
w := s.makeAccountRequest("POST", s.accountPath()+"/copilot/translate", body, map[string]string{})
|
||||
// ShouldBindJSON 解析成功 → handler 继续执行
|
||||
assert.True(s.T(), w.Code == http.StatusOK || w.Code == http.StatusUnprocessableEntity,
|
||||
"expected 200 or 500 for TranslateMessage without content, got %d", w.Code)
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestTranslateMessage_缺少target_language不会返回400() {
|
||||
// ShouldBindJSON 不验证 validate 标签
|
||||
body := map[string]interface{}{
|
||||
"content": "Hello",
|
||||
}
|
||||
w := s.makeAccountRequest("POST", s.accountPath()+"/copilot/translate", body, map[string]string{})
|
||||
assert.True(s.T(), w.Code == http.StatusOK || w.Code == http.StatusUnprocessableEntity,
|
||||
"expected 200 or 500 for TranslateMessage without target_language, got %d", w.Code)
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestTranslateMessage_无效accountID返回400() {
|
||||
body := map[string]interface{}{
|
||||
"content": "Hello",
|
||||
"target_language": "zh",
|
||||
}
|
||||
w := s.makeAccountRequest("POST", "/api/v1/accounts/invalid/copilot/translate", body, map[string]string{})
|
||||
assert.Equal(s.T(), http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
func (s *CopilotThreadHandlerTestSuite) TestTranslateMessage_无效JSON返回400() {
|
||||
w := httptest.NewRecorder()
|
||||
req, _ := http.NewRequest("POST", s.accountPath()+"/copilot/translate", bytes.NewReader([]byte("{invalid}")))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
s.accountRouter.ServeHTTP(w, req)
|
||||
assert.Equal(s.T(), http.StatusBadRequest, w.Code)
|
||||
}
|
||||
|
||||
func TestCopilotThreadHandlerTestSuite(t *testing.T) {
|
||||
suite.Run(t, new(CopilotThreadHandlerTestSuite))
|
||||
}
|
||||
@@ -74,7 +74,7 @@ func (h *SSEStreamHandler) StreamCopilotMessage(c *gin.Context) {
|
||||
c.Header("X-Accel-Buffering", "no") // disable nginx buffering
|
||||
|
||||
// Get thread for context
|
||||
thread, err := h.copilotSvc.GetThread(c.Request.Context(), uint(threadID))
|
||||
thread, err := h.copilotSvc.GetThreadByID(c.Request.Context(), uint(threadID))
|
||||
if err != nil {
|
||||
applogger.L().Errorf("SSE GetThread: %v", err)
|
||||
writeSSEMessage(c, "error", `{"error": "thread not found"}`)
|
||||
|
||||
Reference in New Issue
Block a user