feat(captain): align streaming fallbacks
This commit is contained in:
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user