Fix: shared chatbot completion ignores knowledge base and empty response (#17092)

This commit is contained in:
euvre
2026-07-21 10:43:06 +08:00
committed by GitHub
parent ee8aa03fa0
commit 1c828daea1
3 changed files with 572 additions and 238 deletions

View File

@@ -26,7 +26,9 @@ import (
"context"
"errors"
"fmt"
"hash/fnv"
"strings"
"sync"
"ragflow/internal/agent/canvas"
"ragflow/internal/agent/dsl"
@@ -47,6 +49,15 @@ type BotService struct {
api4ConversationDAO *dao.API4ConversationDAO
agentService *AgentService
llmService *LLMService
pipeline *ChatPipelineService
// persistLocks serialises persistChatbotTurn's read-modify-write
// on a single api_4_conversation row. ChatbotCompletion fetches
// the session before streaming starts, so without this lock two
// concurrent requests on the same session_id would each append
// their turn to the same stale base and the last Update would
// silently drop the other exchange. Striped to a fixed size so
// the lock set does not grow with the number of sessions.
persistLocks [64]sync.Mutex
}
// NewBotService wires a fresh BotService. agentSvc is required for
@@ -59,6 +70,7 @@ func NewBotService(agentSvc *AgentService, llmSvc *LLMService) *BotService {
api4ConversationDAO: dao.NewAPI4ConversationDAO(),
agentService: agentSvc,
llmService: llmSvc,
pipeline: NewChatPipelineService(),
}
}
@@ -190,6 +202,28 @@ type ChatbotCompletionRequest struct {
Question string `json:"question"`
Stream bool `json:"stream"`
Inputs map[string]any `json:"inputs"`
// Quote controls citation generation. Nil means "absent" —
// python bot_api.py defaults it to False for chatbot
// completions, so the service layer mirrors that.
Quote *bool `json:"quote"`
// Reasoning / Internet arrive as bool OR 0/1 number depending
// on the widget; the service layer normalises them before
// handing them to the chat pipeline.
Reasoning any `json:"reasoning"`
Internet any `json:"internet"`
// DocIDs is an optional comma-separated document filter,
// same shape as the regular chat completion kwargs.
DocIDs string `json:"doc_ids"`
}
// persistLock returns the striped mutex guarding one session row's
// read-modify-write in persistChatbotTurn. The modulo runs on the
// unsigned hash — converting to int first would go negative on
// 32-bit architectures and panic with an out-of-bounds index.
func (s *BotService) persistLock(sessionID string) *sync.Mutex {
h := fnv.New32a()
_, _ = h.Write([]byte(sessionID))
return &s.persistLocks[h.Sum32()%uint32(len(s.persistLocks))]
}
// loadCanvas is the IDOR guard for agentbot reads. It mirrors the

View File

@@ -39,6 +39,7 @@ import (
"errors"
"net/http"
"ragflow/internal/utility"
"strings"
"time"
"go.uber.org/zap"
@@ -47,7 +48,6 @@ import (
"ragflow/internal/agent/runtime"
"ragflow/internal/common"
"ragflow/internal/entity"
modelModule "ragflow/internal/entity/models"
)
// ChatbotSSEFrame is one envelope pushed to the SSE writer by the
@@ -68,6 +68,19 @@ type ChatbotSSEFrame struct {
SessionID string `json:"-"`
Done bool `json:"-"`
Err error `json:"-"`
// Final marks the last answer frame of a turn. It is
// rendered as `"final": true` in the data payload so the
// front-end replaces (instead of appends to) the
// accumulated text — required because the final pipeline
// result is decorated (citation markers inserted mid-text)
// and no longer a strict superset of the streamed deltas.
Final bool `json:"-"`
// StartToThink / EndToThink bracket the reasoning segment,
// rendered as start_to_think / end_to_think so the
// front-end can wrap it in <think> markers like the python
// side does.
StartToThink bool `json:"-"`
EndToThink bool `json:"-"`
}
// WriteChatbotFrame emits one python-style SSE frame and flushes the
@@ -101,6 +114,15 @@ func WriteChatbotFrame(w http.ResponseWriter, f ChatbotSSEFrame) error {
"id": nil,
"session_id": f.SessionID,
}
if f.Final {
data["final"] = true
}
if f.StartToThink {
data["start_to_think"] = true
}
if f.EndToThink {
data["end_to_think"] = true
}
// Forward the canvas event type so the front-end can
// distinguish interactive form pauses ("user_inputs",
// "workflow_finished") from plain assistant messages
@@ -262,12 +284,13 @@ func WriteAgentbotFrame(w http.ResponseWriter, f ChatbotSSEFrame) error {
// ChatbotCompletion streams an SSE response for
// /api/v1/chatbots/<dialog_id>/completions.
//
// The full LLM session-lifecycle implementation is added below. It
// is a v1 port: it yields a single frame per turn (the Go LLMBundle
// chat call is non-streaming). A request without session_id only
// creates the prologue-seeded session and streams the prologue back;
// the LLM runs exclusively for follow-up turns that carry a
// session_id.
// The completion runs through ChatPipelineService.AsyncChat — the
// same RAG pipeline that serves the regular chat endpoints — so
// knowledge-base retrieval, the dialog's configured empty_response
// fallback, citations and the system prompt all behave identically
// to the in-app chat. Mirrors the python
// api/db/services/conversation_service.py::async_iframe_completion,
// which delegates to the same async_chat used by regular sessions.
//
// Authorisation: dialog must exist, belong to the requester's tenant,
// and have status == common.StatusDialogValid.
@@ -307,20 +330,13 @@ func (s *BotService) ChatbotCompletion(
// same value, so write/read stay symmetric. We keep this
// behaviour and add the comment so a future reader doesn't
// "fix" it to a tenant-id lookup and break the symmetry.
var session *entity.API4Conversation
if req.SessionID != "" {
session, err = s.api4ConversationDAO.GetBySessionID(req.SessionID, dialogID)
if err != nil {
return nil, common.CodeServerError, err
}
if session == nil || session.UserID != tenantID {
return nil, common.CodeDataError, errors.New("session not found")
}
} else {
// Seed a new session. The Message column is json.RawMessage;
// pre-serialise the prologue turn as a JSON array of
// {role,content,created_at} dicts — same shape the python
// conversation_service.py:253-272 writes. Plan Risk R4.
if req.SessionID == "" {
// No session yet: seed one with the prologue and return it
// immediately WITHOUT running the pipeline. Mirrors python
// async_iframe_completion (conversation_service.py:324-334):
// the share page calls this endpoint once with an empty
// question to obtain a session_id, then sends the real
// questions with that session_id.
prologue := stringFromMap(dialog.PromptConfig, "prologue")
seedMsg, _ := json.Marshal([]map[string]any{
{
@@ -329,7 +345,7 @@ func (s *BotService) ChatbotCompletion(
"created_at": time.Now().Unix(),
},
})
session = &entity.API4Conversation{
session := &entity.API4Conversation{
ID: utility.GenerateUUID(),
DialogID: dialogID,
UserID: tenantID,
@@ -344,8 +360,8 @@ func (s *BotService) ChatbotCompletion(
// session_id is the share page's opening handshake — the
// front-end sends an empty question only to obtain a session.
// Persist the prologue-seeded session and stream the prologue
// back WITHOUT invoking the LLM; running the model here would
// fabricate a reply to a message the user never sent.
// back WITHOUT invoking the pipeline; running the model here
// would fabricate a reply to a message the user never sent.
out := make(chan ChatbotSSEFrame, 2)
go func() {
defer close(out)
@@ -359,135 +375,286 @@ func (s *BotService) ChatbotCompletion(
return out, common.CodeSuccess, nil
}
// 3. Resolve the chat LLM via ModelProviderService. The python
// async_iframe_completion resolves the same way through
// LLMBundle(tenant_id, dialog.llm_id); the Go equivalent is
// GetChatModelConfig → NewChatModel → driver.ChatWithMessages.
//
// If llmService is unwired (test boot path) or the dialog has
// no LLM configured, we surface a sanitized CodeDataError
// rather than echoing the bare error string into the SSE
// envelope — see WriteChatbotFrame's sanitization contract.
session, err := s.api4ConversationDAO.GetBySessionID(req.SessionID, dialogID)
if err != nil {
return nil, common.CodeServerError, err
}
if session == nil || session.UserID != tenantID {
return nil, common.CodeDataError, errors.New("session not found")
}
// 3. Guard rails mirroring the previous implementation: surface
// a sanitised error before any SSE byte is written when the
// service is unwired (test boot path) or the dialog has no LLM
// configured — see WriteChatbotFrame's sanitization contract.
// NewBotService wires both dependencies; the nil checks only
// guard a hand-rolled zero-value BotService against panicking.
if s.llmService == nil {
return nil, common.CodeServerError, errors.New("bot: llm service not wired")
}
if s.pipeline == nil {
return nil, common.CodeServerError, errors.New("bot: chat pipeline not wired")
}
if dialog.LLMID == "" {
return nil, common.CodeDataError, errors.New("no LLM configured for this chatbot")
}
modelProvider := NewModelProviderService()
driver, modelName, apiConfig, _, err := modelProvider.GetChatModelConfig(tenantID, dialog.LLMID)
if err != nil {
return nil, common.CodeDataError, errors.New("no LLM configured for this chatbot")
// 4. Build the pipeline input. The Message column on
// api_4_conversation is a json.RawMessage array of
// {role, content, created_at} dicts; the pipeline expects the
// same filtered shape python builds in async_iframe_completion
// (drop system turns, drop the leading assistant prologue,
// append the new user turn last).
messageID := utility.GenerateUUID()
messages := buildChatbotPipelineMessages(session.Message, req.Question, messageID)
// python bot_api.py:72-73 defaults quote to False for chatbot
// completions when the caller omits it.
kwargs := map[string]interface{}{
"quote": req.Quote != nil && *req.Quote,
}
if reasoning, ok := normalizeBotBoolFlag(req.Reasoning); ok {
kwargs["reasoning"] = reasoning
}
if req.Internet != nil {
kwargs["internet"] = req.Internet
}
if req.DocIDs != "" {
kwargs["doc_ids"] = req.DocIDs
}
chatModel := modelModule.NewChatModel(driver, &modelName, apiConfig)
// 4. Build the prompt from prior conversation history plus the
// new user turn. Without this, a resumed session_id would
// authorise reuse but the LLM call would still be stateless
// turn-to-turn — a Python parity regression for any multi-turn
// chatbot client. The Message column on api_4_conversation is a
// json.RawMessage array of {role, content, created_at} dicts,
// matching the python conversation_service.py:253-272 shape.
messages := historyToMessages(session.Message)
messages = append(messages, modelModule.Message{Role: "user", Content: req.Question})
results, err := s.pipeline.AsyncChat(ctx, tenantID, dialog, messages, true, kwargs)
if err != nil {
return nil, common.CodeServerError, err
}
// 5. Yield frames on a channel.
out := make(chan ChatbotSSEFrame, 4)
// 5. Translate pipeline results into chatbot frames. The
// pipeline streams deltas; the python iframe contract sends the
// full accumulated answer in every frame, so we accumulate here
// (the front-end accepts both shapes, but full-text frames are
// byte-parity with python). The final pipeline result carries
// the decorated full answer plus the retrieval reference.
out := make(chan ChatbotSSEFrame, 16)
go func() {
defer close(out)
resp, callErr := chatModel.ModelDriver.ChatWithMessages(
modelName, messages, chatModel.APIConfig, &modelModule.ChatConfig{}, nil,
var (
fullAnswer string
finalRef map[string]any
errored bool
)
if callErr != nil {
// Log the real error with structured context so
// ops can debug, but do NOT echo the raw
// err.Error() to the client (security M2:
// internal gorm/SQL/file-path leaks).
common.Error("bot: ChatbotCompletion LLM call failed",
callErr,
zap.String("dialog_id", dialogID),
zap.String("session_id", session.ID),
zap.String("llm_id", dialog.LLMID),
)
for res := range results {
if res.Final {
if res.Answer != "" {
// Decorated full answer (citations
// resolved). Replaces the accumulated
// deltas; the empty_response fallback
// path yields an empty final answer, in
// which case the accumulated fallback
// text stands.
fullAnswer = res.Answer
}
if res.Reference != nil {
finalRef = res.Reference
}
if strings.HasPrefix(fullAnswer, "**ERROR**") {
errored = true
}
out <- ChatbotSSEFrame{
Data: fullAnswer,
Reference: referenceOrEmpty(finalRef),
SessionID: session.ID,
Final: true,
}
continue
}
if res.StartToThink || res.EndToThink {
out <- ChatbotSSEFrame{
Data: fullAnswer,
Reference: map[string]any{},
SessionID: session.ID,
StartToThink: res.StartToThink,
EndToThink: res.EndToThink,
}
continue
}
fullAnswer += res.Answer
if len(res.Reference) > 0 {
// The pipeline only populates Reference on the final
// result today; tracking it here keeps finalRef
// correct if a mid-stream result ever carries one.
// Intermediate frames still send an empty reference
// object for wire parity with python
// async_iframe_completion — only the final frame
// carries the retrieval reference.
finalRef = res.Reference
}
out <- ChatbotSSEFrame{
Err: errors.New("an internal error occurred"),
Data: fullAnswer,
Reference: map[string]any{},
SessionID: session.ID,
}
out <- ChatbotSSEFrame{Done: true}
return
}
answer := ""
if resp != nil && resp.Answer != nil {
answer = *resp.Answer
}
// Persist the new turn pair (user + assistant) back to
// api_4_conversation so the NEXT call to ChatbotCompletion
// with the same session_id sees this turn in messages.
// Update errors are logged but do NOT fail the SSE stream
// — the answer has already been produced. The next call
// will rebuild from the prior (pre-this-turn) snapshot,
// losing at most the latest exchange; acceptable for v1.
newTurns := append(historyFromMessages(messages),
map[string]any{"role": "assistant", "content": answer, "created_at": time.Now().Unix()},
)
if updated, mErr := json.Marshal(newTurns); mErr == nil {
session.Message = updated
if uErr := s.api4ConversationDAO.Update(session); uErr != nil {
// 6. Persist the completed turn pair (user + assistant)
// plus the retrieval reference, mirroring python
// API4ConversationService.append_message after the stream.
// Persistence errors are logged but do NOT fail the SSE
// stream — the answer has already been produced. On a
// pipeline-level error ("**ERROR**" answer) nothing is
// persisted, matching the python exception path.
if !errored {
if pErr := s.persistChatbotTurn(session, req.Question, fullAnswer, messageID, finalRef); pErr != nil {
common.Error("bot: ChatbotCompletion session update failed",
uErr,
pErr,
zap.String("dialog_id", dialogID),
zap.String("session_id", session.ID),
)
}
}
out <- ChatbotSSEFrame{
Data: answer,
Reference: map[string]any{},
SessionID: session.ID,
}
out <- ChatbotSSEFrame{Done: true}
}()
return out, common.CodeSuccess, nil
}
// historyToMessages reads the session.Message JSON array of
// {role, content, ...} dicts and projects it onto modelModule.Message
// for the LLM driver. Tolerates an empty / malformed Message column
// by returning an empty slice — the caller appends the new user turn
// so the LLM still receives the current prompt.
func historyToMessages(raw json.RawMessage) []modelModule.Message {
if len(raw) == 0 {
return nil
}
var turns []map[string]any
if err := json.Unmarshal(raw, &turns); err != nil {
return nil
}
out := make([]modelModule.Message, 0, len(turns))
for _, t := range turns {
role, _ := t["role"].(string)
content, _ := t["content"].(string)
if role == "" || content == "" {
// buildChatbotPipelineMessages projects the session.Message JSON
// array plus the new user question onto the message shape
// ChatPipelineService.AsyncChat expects. Mirrors the filtering in
// python async_iframe_completion (conversation_service.py:341-356):
// system turns are dropped and a leading assistant turn (the seeded
// prologue) is dropped so the first message the LLM sees is a user
// turn. Tolerates an empty / malformed Message column by starting
// from just the new question.
func buildChatbotPipelineMessages(raw json.RawMessage, question, messageID string) []map[string]interface{} {
turns := parseChatbotTurns(raw)
turns = append(turns, map[string]any{
"role": "user",
"content": question,
"id": messageID,
})
msg := make([]map[string]interface{}, 0, len(turns))
for _, m := range turns {
role, _ := m["role"].(string)
if role == "system" {
continue
}
out = append(out, modelModule.Message{Role: role, Content: content})
if role == "assistant" && len(msg) == 0 {
continue
}
msg = append(msg, m)
}
return out
return msg
}
// historyFromMessages is the inverse projection — used to write the
// updated turn list back to the api_4_conversation.Message column.
func historyFromMessages(msgs []modelModule.Message) []map[string]any {
out := make([]map[string]any, 0, len(msgs))
now := time.Now().Unix()
for i, m := range msgs {
out = append(out, map[string]any{
"role": m.Role,
"content": m.Content,
"created_at": now + int64(i), // preserve order, monotonic
})
// parseChatbotTurns decodes the session.Message JSON array.
// Returns an empty (non-nil) slice on empty or malformed input so
// callers can always append.
func parseChatbotTurns(raw json.RawMessage) []map[string]any {
turns := make([]map[string]any, 0)
if len(raw) == 0 {
return turns
}
return out
if err := json.Unmarshal(raw, &turns); err != nil || turns == nil {
return make([]map[string]any, 0)
}
return turns
}
// persistChatbotTurn appends the finished user/assistant turn pair
// and the retrieval reference to the api_4_conversation row so the
// next ChatbotCompletion call with the same session_id sees this
// turn in its history. Mirrors python
// API4ConversationService.append_message.
func (s *BotService) persistChatbotTurn(
session *entity.API4Conversation, question, answer, messageID string, reference map[string]any,
) error {
// Serialise the read-modify-write per session and re-read the row
// inside the lock: the caller's session was loaded before the
// stream ran, so a concurrent request on the same session_id may
// already have appended its own turn. Without the lock + re-read
// the last Update would silently drop the other exchange.
lock := s.persistLock(session.ID)
lock.Lock()
defer lock.Unlock()
fresh, err := s.api4ConversationDAO.GetBySessionID(session.ID, session.DialogID)
if err != nil {
return err
}
if fresh != nil {
session = fresh
}
now := time.Now().Unix()
turns := parseChatbotTurns(session.Message)
// Both turns of the pair share messageID by design: a Q&A exchange
// is addressed as a unit — mirrors the in-app chat convention
// where the answer id is derived from the question id so the pair
// is deleted together (web/src/hooks/logic-hooks.ts
// buildMessageUuid).
turns = append(turns,
map[string]any{
"role": "user",
"content": question,
"id": messageID,
"created_at": now,
},
map[string]any{
"role": "assistant",
"content": answer,
"id": messageID,
"created_at": now,
},
)
rawMsg, err := json.Marshal(turns)
if err != nil {
return err
}
session.Message = rawMsg
refs := make([]any, 0)
if len(session.Reference) > 0 {
// Tolerate malformed / missing reference history — the
// message turns above are the authoritative history;
// a lost reference list degrades citation display only.
_ = json.Unmarshal(session.Reference, &refs)
}
if reference == nil {
reference = map[string]any{"chunks": []any{}, "doc_aggs": []any{}}
}
refs = append(refs, reference)
rawRef, err := json.Marshal(refs)
if err != nil {
return err
}
session.Reference = rawRef
return s.api4ConversationDAO.Update(session)
}
// normalizeBotBoolFlag coerces the JSON-encoded reasoning / internet
// flags to a bool. Widgets send them as true/false or 0/1 numbers;
// ok=false means the value was absent or unrecognised and the caller
// should leave the pipeline default in place.
func normalizeBotBoolFlag(v any) (value, ok bool) {
switch x := v.(type) {
case bool:
return x, true
case float64:
if x == 0 || x == 1 {
return x == 1, true
}
case int:
if x == 0 || x == 1 {
return x == 1, true
}
}
return false, false
}
// referenceOrEmpty returns ref or an empty map so SSE frames always
// carry a JSON object in the reference field, never null.
func referenceOrEmpty(ref map[string]any) map[string]any {
if ref == nil {
return map[string]any{}
}
return ref
}

View File

@@ -14,11 +14,13 @@
// limitations under the License.
//
// Tests for the conversation-history round-trip helpers used by
// BotService.ChatbotCompletion. Locks in review Finding 8 — a resumed
// session_id must carry prior turns (assistant prologue + earlier
// user/assistant exchanges) into the next LLM call so multi-turn
// chatbot clients retain context.
// Tests for the message-building / flag-normalisation helpers used
// by BotService.ChatbotCompletion. A resumed session_id must carry
// prior turns (assistant prologue + earlier user/assistant
// exchanges) into the next pipeline call so multi-turn chatbot
// clients retain context, and the filtering must match python
// async_iframe_completion (drop system turns and the leading
// assistant prologue).
package service
@@ -27,148 +29,94 @@ import (
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"strings"
"sync"
"testing"
"ragflow/internal/agent/canvas"
"ragflow/internal/common"
"ragflow/internal/dao"
"ragflow/internal/entity"
modelModule "ragflow/internal/entity/models"
)
func TestHistoryToMessages_Empty(t *testing.T) {
// A freshly-seeded session with no prior turns returns an empty
// slice. Caller appends the new user turn; LLM receives only
// the current prompt. Matches python conversation_service seed.
got := historyToMessages(nil)
if len(got) != 0 {
t.Fatalf("nil raw: want 0 messages, got %d", len(got))
func TestBuildChatbotPipelineMessages_Empty(t *testing.T) {
// A freshly-seeded session with no prior turns produces just
// the new user question. Matches python conversation_service.
msgs := buildChatbotPipelineMessages(nil, "hi", "msg-1")
if len(msgs) != 1 {
t.Fatalf("nil raw: want 1 message, got %d", len(msgs))
}
got = historyToMessages(json.RawMessage(`[]`))
if len(got) != 0 {
t.Fatalf("empty array: want 0 messages, got %d", len(got))
if msgs[0]["role"] != "user" || msgs[0]["content"] != "hi" || msgs[0]["id"] != "msg-1" {
t.Errorf("got %+v", msgs[0])
}
msgs = buildChatbotPipelineMessages(json.RawMessage(`[]`), "hi", "msg-1")
if len(msgs) != 1 {
t.Fatalf("empty array: want 1 message, got %d", len(msgs))
}
}
func TestHistoryToMessages_RoundTrip(t *testing.T) {
// Simulate a session with: 1 prologue assistant turn + 1 prior
// user/assistant pair. The LLM must see all 3 prior turns
// before the new user turn is appended.
func TestBuildChatbotPipelineMessages_DropsLeadingAssistantAndSystem(t *testing.T) {
// Prologue (leading assistant) and system turns must not reach
// the pipeline; later assistant turns are kept so the LLM sees
// prior exchanges.
turns := []map[string]any{
{"role": "assistant", "content": "Hello, how can I help?", "created_at": 1},
{"role": "user", "content": "What is Go?", "created_at": 2},
{"role": "assistant", "content": "Go is a compiled language.", "created_at": 3},
{"role": "system", "content": "hidden", "created_at": 2},
{"role": "user", "content": "What is Go?", "created_at": 3},
{"role": "assistant", "content": "Go is a compiled language.", "created_at": 4},
}
raw, err := json.Marshal(turns)
if err != nil {
t.Fatalf("marshal seed: %v", err)
}
msgs := historyToMessages(raw)
msgs := buildChatbotPipelineMessages(raw, "and Rust?", "msg-2")
if len(msgs) != 3 {
t.Fatalf("want 3 prior messages, got %d", len(msgs))
t.Fatalf("want 3 messages (user, assistant, new user), got %d: %+v", len(msgs), msgs)
}
if msgs[0].Role != "assistant" || msgs[0].Content != "Hello, how can I help?" {
t.Errorf("turn 0: role=%q content=%q", msgs[0].Role, msgs[0].Content)
if msgs[0]["role"] != "user" || msgs[0]["content"] != "What is Go?" {
t.Errorf("turn 0: %+v", msgs[0])
}
if msgs[1].Role != "user" || msgs[1].Content != "What is Go?" {
t.Errorf("turn 1: role=%q content=%q", msgs[1].Role, msgs[1].Content)
if msgs[1]["role"] != "assistant" || msgs[1]["content"] != "Go is a compiled language." {
t.Errorf("turn 1: %+v", msgs[1])
}
if msgs[2].Role != "assistant" || msgs[2].Content != "Go is a compiled language." {
t.Errorf("turn 2: role=%q content=%q", msgs[2].Role, msgs[2].Content)
if msgs[2]["role"] != "user" || msgs[2]["content"] != "and Rust?" || msgs[2]["id"] != "msg-2" {
t.Errorf("turn 2: %+v", msgs[2])
}
}
func TestHistoryToMessages_Malformed(t *testing.T) {
// Malformed JSON must not panic; returns nil so caller falls back
// to a fresh single-turn LLM call rather than failing the request.
got := historyToMessages(json.RawMessage(`not json`))
if got != nil {
t.Fatalf("malformed raw: want nil, got %v", got)
func TestBuildChatbotPipelineMessages_Malformed(t *testing.T) {
// Malformed JSON must not panic; falls back to just the new
// question rather than failing the request.
msgs := buildChatbotPipelineMessages(json.RawMessage(`not json`), "hi", "msg-3")
if len(msgs) != 1 || msgs[0]["content"] != "hi" {
t.Fatalf("malformed raw: want single user turn, got %+v", msgs)
}
}
func TestHistoryToMessages_SkipsEmptyFields(t *testing.T) {
// Defensive: turns missing role or content are dropped, not
// passed to the LLM as empty messages.
turns := []map[string]any{
{"role": "assistant", "content": "valid", "created_at": 1},
{"role": "", "content": "no role", "created_at": 2},
{"role": "user", "content": "", "created_at": 3},
{"role": "user", "content": "second valid", "created_at": 4},
func TestNormalizeBotBoolFlag(t *testing.T) {
cases := []struct {
in any
value bool
ok bool
}{
{true, true, true},
{false, false, true},
{float64(1), true, true},
{float64(0), false, true},
{1, true, true},
{0, false, true},
{nil, false, false},
{"yes", false, false},
{float64(2), false, false},
}
raw, _ := json.Marshal(turns)
msgs := historyToMessages(raw)
if len(msgs) != 2 {
t.Fatalf("want 2 valid turns, got %d", len(msgs))
}
if msgs[0].Content != "valid" || msgs[1].Content != "second valid" {
t.Errorf("got %+v", msgs)
}
}
func TestHistoryFromMessages_PreservesOrder(t *testing.T) {
// The LLM driver returns messages in the same order the input
// was provided. The round-trip must preserve that order so the
// next call to ChatbotCompletion sees a coherent history.
msgs := []modelModule.Message{
{Role: "assistant", Content: "first"},
{Role: "user", Content: "second"},
{Role: "assistant", Content: "third"},
}
turns := historyFromMessages(msgs)
if len(turns) != 3 {
t.Fatalf("want 3 turns, got %d", len(turns))
}
for i, want := range []string{"first", "second", "third"} {
if turns[i]["content"] != want {
t.Errorf("turn %d content = %v, want %q", i, turns[i]["content"], want)
}
if turns[i]["role"] != msgs[i].Role {
t.Errorf("turn %d role = %v, want %q", i, turns[i]["role"], msgs[i].Role)
}
}
}
func TestHistoryRoundTrip_PreservesPriorTurns(t *testing.T) {
// End-to-end: prior JSON → history → back to JSON must be
// semantically identical (modulo the created_at monotonic
// adjustment that historyFromMessages applies for ordering).
turns := []map[string]any{
{"role": "assistant", "content": "p1", "created_at": int64(100)},
{"role": "user", "content": "p2", "created_at": int64(200)},
}
raw, _ := json.Marshal(turns)
msgs := historyToMessages(raw)
// Caller appends a new user turn (the current request).
msgs = append(msgs, modelModule.Message{Role: "user", Content: "current"})
// Round-trip back to JSON for storage.
newTurns := historyFromMessages(msgs)
raw2, err := json.Marshal(newTurns)
if err != nil {
t.Fatalf("marshal round-trip: %v", err)
}
var got []map[string]any
if err := json.Unmarshal(raw2, &got); err != nil {
t.Fatalf("unmarshal round-trip: %v", err)
}
if len(got) != 3 {
t.Fatalf("want 3 turns after round-trip, got %d", len(got))
}
expected := []struct{ role, content string }{
{"assistant", "p1"},
{"user", "p2"},
{"user", "current"},
}
for i, want := range expected {
if got[i]["role"] != want.role || got[i]["content"] != want.content {
t.Errorf("turn %d: got role=%v content=%v, want role=%q content=%q",
i, got[i]["role"], got[i]["content"], want.role, want.content)
for _, c := range cases {
value, ok := normalizeBotBoolFlag(c.in)
if value != c.value || ok != c.ok {
t.Errorf("normalizeBotBoolFlag(%v) = (%v, %v), want (%v, %v)",
c.in, value, ok, c.value, c.ok)
}
}
}
@@ -391,12 +339,197 @@ func TestBotService_ChatbotCompletion_NewSessionSkipsLLM(t *testing.T) {
if derr != nil || row == nil {
t.Fatalf("session not persisted: row=%v err=%v", row, derr)
}
msgs := historyToMessages(row.Message)
if len(msgs) != 1 || msgs[0].Role != "assistant" || msgs[0].Content != prologue {
msgs := parseChatbotTurns(row.Message)
if len(msgs) != 1 || msgs[0]["role"] != "assistant" || msgs[0]["content"] != prologue {
t.Errorf("persisted messages = %+v, want single prologue assistant turn", msgs)
}
}
func TestParseChatbotTurns(t *testing.T) {
// Empty / malformed input must yield an empty (non-nil) slice so
// callers can always append.
for _, raw := range []json.RawMessage{
nil,
json.RawMessage(`[]`),
json.RawMessage(`null`),
json.RawMessage(`not json`),
} {
turns := parseChatbotTurns(raw)
if turns == nil {
t.Errorf("parseChatbotTurns(%q) returned a nil slice", string(raw))
}
if len(turns) != 0 {
t.Errorf("parseChatbotTurns(%q) = %+v, want empty", string(raw), turns)
}
}
turns := parseChatbotTurns(json.RawMessage(`[{"role":"user","content":"hi"}]`))
if len(turns) != 1 || turns[0]["role"] != "user" || turns[0]["content"] != "hi" {
t.Errorf("valid input: got %+v", turns)
}
}
func TestReferenceOrEmpty(t *testing.T) {
got := referenceOrEmpty(nil)
if got == nil || len(got) != 0 {
t.Errorf("nil reference must yield an empty non-nil map, got %v", got)
}
ref := map[string]any{"chunks": []any{map[string]any{"chunk_id": "c1"}}}
got = referenceOrEmpty(ref)
chunks, _ := got["chunks"].([]any)
if len(chunks) != 1 {
t.Errorf("non-nil reference must pass through, got %+v", got)
}
}
// TestPersistChatbotTurn_AppendsPairAndReference covers the turn
// persistence after a completed stream: the user/assistant pair is
// appended to the existing history, shares one message id, and the
// retrieval reference is appended to the reference list.
func TestPersistChatbotTurn_AppendsPairAndReference(t *testing.T) {
db := setupServiceTestDB(t)
pushServiceDB(t, db)
seed, _ := json.Marshal([]map[string]any{
{"role": "assistant", "content": "Hello!", "created_at": 1},
})
sess := &entity.API4Conversation{
ID: "sess-p1",
DialogID: "dlg-p1",
UserID: "tenant-1",
Message: seed,
}
if err := dao.NewAPI4ConversationDAO().Create(sess); err != nil {
t.Fatalf("seed session: %v", err)
}
svc := NewBotService(nil, nil)
ref := map[string]any{
"chunks": []any{map[string]any{"chunk_id": "c1"}},
"doc_aggs": []any{},
}
if err := svc.persistChatbotTurn(sess, "What is Go?", "A language.", "msg-p1", ref); err != nil {
t.Fatalf("persistChatbotTurn: %v", err)
}
row, err := dao.NewAPI4ConversationDAO().GetBySessionID("sess-p1", "dlg-p1")
if err != nil || row == nil {
t.Fatalf("re-read session: row=%v err=%v", row, err)
}
turns := parseChatbotTurns(row.Message)
if len(turns) != 3 {
t.Fatalf("want 3 turns (prologue + pair), got %d: %+v", len(turns), turns)
}
// The Q&A pair shares the message id so the exchange is addressed
// as a unit, matching the in-app chat pairing convention.
if turns[1]["role"] != "user" || turns[1]["content"] != "What is Go?" || turns[1]["id"] != "msg-p1" {
t.Errorf("user turn: %+v", turns[1])
}
if turns[2]["role"] != "assistant" || turns[2]["content"] != "A language." || turns[2]["id"] != "msg-p1" {
t.Errorf("assistant turn: %+v", turns[2])
}
var refs []map[string]any
if err := json.Unmarshal(row.Reference, &refs); err != nil {
t.Fatalf("reference decode: %v", err)
}
if len(refs) != 1 {
t.Fatalf("want 1 reference entry, got %d: %+v", len(refs), refs)
}
chunks, _ := refs[0]["chunks"].([]any)
if len(chunks) != 1 {
t.Errorf("reference chunks: %+v", refs[0])
}
}
// TestPersistChatbotTurn_NilReferenceDefaultsToEmpty ensures a turn
// without retrieval results still records an empty reference object,
// keeping the reference list index-aligned with the turn pairs.
func TestPersistChatbotTurn_NilReferenceDefaultsToEmpty(t *testing.T) {
db := setupServiceTestDB(t)
pushServiceDB(t, db)
sess := &entity.API4Conversation{ID: "sess-p2", DialogID: "dlg-p1", UserID: "tenant-1"}
if err := dao.NewAPI4ConversationDAO().Create(sess); err != nil {
t.Fatalf("seed session: %v", err)
}
svc := NewBotService(nil, nil)
if err := svc.persistChatbotTurn(sess, "q", "a", "msg-p2", nil); err != nil {
t.Fatalf("persistChatbotTurn: %v", err)
}
row, err := dao.NewAPI4ConversationDAO().GetBySessionID("sess-p2", "dlg-p1")
if err != nil || row == nil {
t.Fatalf("re-read session: row=%v err=%v", row, err)
}
var refs []map[string]any
if err := json.Unmarshal(row.Reference, &refs); err != nil {
t.Fatalf("reference decode: %v", err)
}
if len(refs) != 1 {
t.Fatalf("want 1 reference entry, got %+v", refs)
}
if chunks, ok := refs[0]["chunks"].([]any); !ok || len(chunks) != 0 {
t.Errorf("default reference must carry empty chunks: %+v", refs[0])
}
if aggs, ok := refs[0]["doc_aggs"].([]any); !ok || len(aggs) != 0 {
t.Errorf("default reference must carry empty doc_aggs: %+v", refs[0])
}
}
// TestPersistChatbotTurn_ConcurrentSameSession guards the
// read-modify-write race on session.Message: two requests persisting
// to the same session_id at the same time must both land — without
// the persist lock + re-read the last Update would silently drop one
// exchange.
func TestPersistChatbotTurn_ConcurrentSameSession(t *testing.T) {
db := setupServiceTestDB(t)
pushServiceDB(t, db)
// In-memory sqlite opens a fresh database per connection, so pin
// the pool to a single connection for the concurrent goroutines.
sqlDB, err := db.DB()
if err != nil {
t.Fatalf("db handle: %v", err)
}
sqlDB.SetMaxOpenConns(1)
sess := &entity.API4Conversation{ID: "sess-p3", DialogID: "dlg-p1", UserID: "tenant-1"}
if err := dao.NewAPI4ConversationDAO().Create(sess); err != nil {
t.Fatalf("seed session: %v", err)
}
svc := NewBotService(nil, nil)
// Both callers hold the same stale snapshot loaded before their
// streams started — exactly the race the lock fixes.
var wg sync.WaitGroup
for i, turn := range [][2]string{{"q1", "a1"}, {"q2", "a2"}} {
wg.Add(1)
go func(i int, q, a string) {
defer wg.Done()
if err := svc.persistChatbotTurn(sess, q, a, fmt.Sprintf("msg-p3-%d", i), nil); err != nil {
t.Errorf("persistChatbotTurn %d: %v", i, err)
}
}(i, turn[0], turn[1])
}
wg.Wait()
row, err := dao.NewAPI4ConversationDAO().GetBySessionID("sess-p3", "dlg-p1")
if err != nil || row == nil {
t.Fatalf("re-read session: row=%v err=%v", row, err)
}
if turns := parseChatbotTurns(row.Message); len(turns) != 4 {
t.Fatalf("want both exchanges persisted (4 turns), got %d: %+v", len(turns), turns)
}
var refs []map[string]any
if err := json.Unmarshal(row.Reference, &refs); err != nil {
t.Fatalf("reference decode: %v", err)
}
if len(refs) != 2 {
t.Fatalf("want 2 reference entries, got %d: %+v", len(refs), refs)
}
}
// recordingResponseWriter is a minimal http.ResponseWriter stub
// for SSE frame tests. Tracks writes so the test can assert the
// emitted frame contents.