feat(captain): align streaming fallbacks

This commit is contained in:
2026-06-05 13:57:30 +08:00
parent a7cd53aa62
commit 72211e6dda
8 changed files with 329 additions and 67 deletions
+20 -14
View File
@@ -111,13 +111,13 @@ func (h *CaptainTaskHandler) Rewrite(c *gin.Context) {
func (h *CaptainTaskHandler) StreamReplySuggestion(c *gin.Context) {
accountID, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
captainWriteSSEError(c, "invalid account_id")
captainWriteSSEError(c, http.StatusBadRequest, "invalid account_id")
return
}
var req service.TaskReplySuggestionRequest
if err := c.ShouldBindJSON(&req); err != nil {
captainWriteSSEError(c, "invalid request body: "+err.Error())
captainWriteSSEError(c, http.StatusBadRequest, "invalid request body: "+err.Error())
return
}
@@ -144,8 +144,7 @@ func (h *CaptainTaskHandler) StreamReplySuggestion(c *gin.Context) {
if err != nil && err != io.EOF {
applogger.L().Errorf("StreamReplySuggestion: %v", err)
captainWriteSSEMessage(c, "error", fmt.Sprintf(`{"error": "%s"}`, captainEscapeJSONString(err.Error())))
captainWriteSSEMessage(c, "done", `{"done": true}`)
captainWriteSSETaskError(c, err)
}
}
@@ -154,13 +153,13 @@ func (h *CaptainTaskHandler) StreamReplySuggestion(c *gin.Context) {
func (h *CaptainTaskHandler) StreamSummarize(c *gin.Context) {
accountID, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
captainWriteSSEError(c, "invalid account_id")
captainWriteSSEError(c, http.StatusBadRequest, "invalid account_id")
return
}
var req service.TaskSummarizeRequest
if err := c.ShouldBindJSON(&req); err != nil {
captainWriteSSEError(c, "invalid request body: "+err.Error())
captainWriteSSEError(c, http.StatusBadRequest, "invalid request body: "+err.Error())
return
}
@@ -187,8 +186,7 @@ func (h *CaptainTaskHandler) StreamSummarize(c *gin.Context) {
if err != nil && err != io.EOF {
applogger.L().Errorf("StreamSummarize: %v", err)
captainWriteSSEMessage(c, "error", fmt.Sprintf(`{"error": "%s"}`, captainEscapeJSONString(err.Error())))
captainWriteSSEMessage(c, "done", `{"done": true}`)
captainWriteSSETaskError(c, err)
}
}
@@ -197,13 +195,13 @@ func (h *CaptainTaskHandler) StreamSummarize(c *gin.Context) {
func (h *CaptainTaskHandler) StreamRewrite(c *gin.Context) {
accountID, err := strconv.ParseUint(c.Param("id"), 10, 64)
if err != nil {
captainWriteSSEError(c, "invalid account_id")
captainWriteSSEError(c, http.StatusBadRequest, "invalid account_id")
return
}
var req service.TaskRewriteRequest
if err := c.ShouldBindJSON(&req); err != nil {
captainWriteSSEError(c, "invalid request body: "+err.Error())
captainWriteSSEError(c, http.StatusBadRequest, "invalid request body: "+err.Error())
return
}
@@ -230,8 +228,7 @@ func (h *CaptainTaskHandler) StreamRewrite(c *gin.Context) {
if err != nil && err != io.EOF {
applogger.L().Errorf("StreamRewrite: %v", err)
captainWriteSSEMessage(c, "error", fmt.Sprintf(`{"error": "%s"}`, captainEscapeJSONString(err.Error())))
captainWriteSSEMessage(c, "done", `{"done": true}`)
captainWriteSSETaskError(c, err)
}
}
@@ -253,12 +250,21 @@ func captainWriteSSEMessage(c *gin.Context, event string, data string) {
}
// captainWriteSSEError writes an error SSE event followed by a done event.
func captainWriteSSEError(c *gin.Context, errMsg string) {
func captainWriteSSEError(c *gin.Context, status int, errMsg string) {
captainSetSSEHeaders(c)
captainWriteSSEMessage(c, "error", fmt.Sprintf(`{"error": "%s"}`, captainEscapeJSONString(errMsg)))
captainWriteSSEMessage(c, "error", fmt.Sprintf(`{"error": "%s", "status": %d, "done": true}`, captainEscapeJSONString(errMsg), status))
captainWriteSSEMessage(c, "done", `{"done": true}`)
}
func captainWriteSSETaskError(c *gin.Context, err error) {
status, message, ok := service.CaptainTaskErrorStatus(err)
if !ok {
status = http.StatusUnprocessableEntity
message = err.Error()
}
captainWriteSSEError(c, status, message)
}
// captainEscapeJSONString escapes special characters for safe JSON embedding.
func captainEscapeJSONString(s string) string {
var result strings.Builder
@@ -7,6 +7,7 @@ import (
"fmt"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/gin-gonic/gin"
@@ -21,8 +22,10 @@ import (
)
type mockCaptainTaskHandlerLLM struct {
response *llm.ChatResponse
err error
response *llm.ChatResponse
err error
streamChunks []llm.StreamChunk
streamErr error
}
func (m *mockCaptainTaskHandlerLLM) ChatCompletion(_ context.Context, _ llm.ChatRequest) (*llm.ChatResponse, error) {
@@ -31,7 +34,16 @@ func (m *mockCaptainTaskHandlerLLM) ChatCompletion(_ context.Context, _ llm.Chat
func (m *mockCaptainTaskHandlerLLM) CreateEmbedding(_ context.Context, _ llm.EmbeddingRequest) (*llm.EmbeddingResponse, error) {
return nil, nil
}
func (m *mockCaptainTaskHandlerLLM) ChatCompletionStream(_ context.Context, _ llm.ChatRequest, _ func(llm.StreamChunk) error) error {
func (m *mockCaptainTaskHandlerLLM) ChatCompletionStream(_ context.Context, _ llm.ChatRequest, onChunk func(llm.StreamChunk) error) error {
if m.streamErr != nil {
return m.streamErr
}
for _, chunk := range m.streamChunks {
if err := onChunk(chunk); err != nil {
return err
}
}
return nil
}
@@ -101,3 +113,67 @@ func TestCaptainTaskHandler_Rewrite_NoProviderRawDisabled(t *testing.T) {
assert.Equal(t, "Captain is disabled", resp["error"])
assert.NotContains(t, resp, "success")
}
func TestCaptainTaskHandler_StreamRewrite_NoProviderDisabledSSE(t *testing.T) {
handler, _ := setupCaptainTaskHandlerTest(t, nil)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Params = gin.Params{{Key: "id", Value: "1"}}
c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/accounts/1/captain/tasks/rewrite/stream", bytes.NewReader([]byte(`{"content":"hello","operation":"professional"}`)))
c.Request.Header.Set("Content-Type", "application/json")
handler.StreamRewrite(c)
body := w.Body.String()
assert.Equal(t, "text/event-stream", w.Header().Get("Content-Type"))
assert.Contains(t, body, "event: error")
assert.Contains(t, body, `"error": "Captain is disabled"`)
assert.Contains(t, body, `"status": 422`)
assert.Contains(t, body, "event: done")
}
func TestCaptainTaskHandler_StreamSummarize_ChatwootDisplayID(t *testing.T) {
provider := &mockCaptainTaskHandlerLLM{streamChunks: []llm.StreamChunk{
{Choices: []llm.StreamChoice{{Delta: llm.StreamDelta{Content: "Short"}}}},
{Choices: []llm.StreamChoice{{Delta: llm.StreamDelta{Content: " summary"}, FinishReason: "stop"}}},
}}
handler, db := setupCaptainTaskHandlerTest(t, provider)
displayID := uint(456)
conv := &model.Conversation{AccountID: 1, DisplayID: &displayID, Status: "open", ChannelType: "web_widget", Channel: "web_widget"}
require.NoError(t, db.Create(conv).Error)
require.NoError(t, db.Create(&model.Message{ConversationID: conv.ID, AccountID: 1, SenderType: "contact", MessageType: "incoming", Content: "Need help"}).Error)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Params = gin.Params{{Key: "id", Value: "1"}}
c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/accounts/1/captain/tasks/summarize/stream", bytes.NewReader([]byte(`{"conversation_display_id":456}`)))
c.Request.Header.Set("Content-Type", "application/json")
handler.StreamSummarize(c)
body := w.Body.String()
assert.Contains(t, body, "event: message")
assert.Contains(t, body, `"content": "Short"`)
assert.Contains(t, body, `"content": " summary"`)
assert.Contains(t, body, "event: done")
assert.NotContains(t, body, "success")
}
func TestCaptainTaskHandler_StreamRewrite_InvalidOperationSSE(t *testing.T) {
handler, _ := setupCaptainTaskHandlerTest(t, &mockCaptainTaskHandlerLLM{})
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Params = gin.Params{{Key: "id", Value: "1"}}
c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/accounts/1/captain/tasks/rewrite/stream", strings.NewReader(`{"content":"hello","operation":"pirate"}`))
c.Request.Header.Set("Content-Type", "application/json")
handler.StreamRewrite(c)
body := w.Body.String()
assert.Contains(t, body, "event: error")
assert.Contains(t, body, `"error": "Invalid operation: pirate"`)
assert.Contains(t, body, `"status": 422`)
assert.Contains(t, body, "event: done")
}
+21 -1
View File
@@ -349,17 +349,37 @@ func copilotThreadPayload(thread *model.CopilotThread) gin.H {
}
}
func copilotThreadPushPayload(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),
"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),
"copilot_thread": copilotThreadPushPayload(&message.CopilotThread),
"account_id": message.AccountID,
}
}
func copilotMessagePushPayload(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": copilotThreadPushPayload(&message.CopilotThread),
}
}
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"}
@@ -237,6 +237,48 @@ func TestCopilotThreadMessagesListAndCreateUseNestedPayloads(t *testing.T) {
require.Len(t, messages, 4)
}
func TestCopilotMessagePayloadUsesThreadPushShape(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/")
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)
message := messages[0].(map[string]any)
nestedThread := message["copilot_thread"].(map[string]any)
require.Equal(t, thread["id"], nestedThread["id"])
require.NotNil(t, nestedThread["user"])
require.Equal(t, thread["account_id"], nestedThread["account_id"])
require.Nil(t, nestedThread["assistant"])
require.NotNil(t, message["account_id"])
}
func TestCopilotMessagePushPayloadMatchesChatwootEventData(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/")
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)
var stored model.CopilotMessage
require.NoError(t, f.db.Preload("CopilotThread.User").Preload("CopilotThread.Assistant").First(&stored, uint(created["id"].(float64))).Error)
push := copilotMessagePushPayload(&stored)
nestedThread := push["copilot_thread"].(gin.H)
require.Equal(t, created["id"], float64(push["id"].(uint)))
require.Equal(t, "user", string(push["message_type"].(model.CopilotMessageType)))
require.Equal(t, "Follow up", push["message"].(map[string]any)["content"])
require.Nil(t, push["account_id"])
require.Equal(t, stored.CopilotThreadID, nestedThread["id"])
require.Nil(t, nestedThread["assistant"])
require.NotNil(t, nestedThread["user"])
}
func TestCopilotThreadMessagesAreAccountAndUserScoped(t *testing.T) {
f := newCopilotParityFixture(t)
thread := f.createThread(t, "Private thread")
@@ -87,6 +87,11 @@ func (h *SSEStreamHandler) StreamCopilotMessage(c *gin.Context) {
writeSSEMessage(c, "done", `{"done": true}`)
return
}
if h.llmProvider == nil {
writeSSEMessage(c, "error", `{"error": "Captain is disabled", "status": 422, "done": true}`)
writeSSEMessage(c, "done", `{"done": true}`)
return
}
// Build chat messages from thread history + new user message
chatMessages := buildStreamChatMessages(thread, req.Content)