fix(seahorse): repair reasoning_content on bootstrap prefix
This commit is contained in:
parent
89e875bd7b
commit
8467125207
2 changed files with 79 additions and 9 deletions
|
|
@ -434,7 +434,7 @@ func (e *Engine) Bootstrap(ctx context.Context, sessionKey string, messages []Me
|
||||||
// Fast path: DB has same count and exact match → no-op
|
// Fast path: DB has same count and exact match → no-op
|
||||||
if len(dbMsgs) == len(messages) {
|
if len(dbMsgs) == len(messages) {
|
||||||
matched := true
|
matched := true
|
||||||
for i := 0; i < len(messages); i++ {
|
for i := range messages {
|
||||||
if !messageMatches(dbMsgs[i], messages[i]) {
|
if !messageMatches(dbMsgs[i], messages[i]) {
|
||||||
matched = false
|
matched = false
|
||||||
break
|
break
|
||||||
|
|
@ -451,18 +451,15 @@ func (e *Engine) Bootstrap(ctx context.Context, sessionKey string, messages []Me
|
||||||
// summaries/context behind after a partial raw-message rebuild.
|
// summaries/context behind after a partial raw-message rebuild.
|
||||||
if repaired, err := e.repairBootstrapReasoningContent(ctx, dbMsgs, messages); err != nil {
|
if repaired, err := e.repairBootstrapReasoningContent(ctx, dbMsgs, messages); err != nil {
|
||||||
return fmt.Errorf("bootstrap: repair reasoning_content: %w", err)
|
return fmt.Errorf("bootstrap: repair reasoning_content: %w", err)
|
||||||
} else if repaired {
|
} else if repaired && len(dbMsgs) == len(messages) {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Find longest matching prefix from the start
|
// Find longest matching prefix from the start
|
||||||
anchor := -1
|
anchor := -1
|
||||||
compareLen := len(dbMsgs)
|
compareLen := min(len(dbMsgs), len(messages))
|
||||||
if compareLen > len(messages) {
|
|
||||||
compareLen = len(messages)
|
|
||||||
}
|
|
||||||
|
|
||||||
for i := 0; i < compareLen; i++ {
|
for i := range compareLen {
|
||||||
if messageMatches(dbMsgs[i], messages[i]) {
|
if messageMatches(dbMsgs[i], messages[i]) {
|
||||||
anchor = i
|
anchor = i
|
||||||
} else {
|
} else {
|
||||||
|
|
@ -549,16 +546,19 @@ func (e *Engine) Bootstrap(ctx context.Context, sessionKey string, messages []Me
|
||||||
}
|
}
|
||||||
|
|
||||||
func (e *Engine) repairBootstrapReasoningContent(ctx context.Context, dbMsgs, messages []Message) (bool, error) {
|
func (e *Engine) repairBootstrapReasoningContent(ctx context.Context, dbMsgs, messages []Message) (bool, error) {
|
||||||
if len(dbMsgs) != len(messages) || len(dbMsgs) == 0 {
|
if len(dbMsgs) == 0 || len(messages) == 0 {
|
||||||
return false, nil
|
return false, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
overlap := min(len(messages), len(dbMsgs))
|
||||||
|
|
||||||
var updates []struct {
|
var updates []struct {
|
||||||
|
index int
|
||||||
messageID int64
|
messageID int64
|
||||||
reasoningContent string
|
reasoningContent string
|
||||||
}
|
}
|
||||||
|
|
||||||
for i := range messages {
|
for i := range overlap {
|
||||||
if !messageMatchesIgnoringReasoning(dbMsgs[i], messages[i]) {
|
if !messageMatchesIgnoringReasoning(dbMsgs[i], messages[i]) {
|
||||||
return false, nil
|
return false, nil
|
||||||
}
|
}
|
||||||
|
|
@ -569,9 +569,11 @@ func (e *Engine) repairBootstrapReasoningContent(ctx context.Context, dbMsgs, me
|
||||||
return false, nil
|
return false, nil
|
||||||
}
|
}
|
||||||
updates = append(updates, struct {
|
updates = append(updates, struct {
|
||||||
|
index int
|
||||||
messageID int64
|
messageID int64
|
||||||
reasoningContent string
|
reasoningContent string
|
||||||
}{
|
}{
|
||||||
|
index: i,
|
||||||
messageID: dbMsgs[i].ID,
|
messageID: dbMsgs[i].ID,
|
||||||
reasoningContent: messages[i].ReasoningContent,
|
reasoningContent: messages[i].ReasoningContent,
|
||||||
})
|
})
|
||||||
|
|
@ -585,6 +587,7 @@ func (e *Engine) repairBootstrapReasoningContent(ctx context.Context, dbMsgs, me
|
||||||
if err := e.store.UpdateMessageReasoningContent(ctx, update.messageID, update.reasoningContent); err != nil {
|
if err := e.store.UpdateMessageReasoningContent(ctx, update.messageID, update.reasoningContent); err != nil {
|
||||||
return false, err
|
return false, err
|
||||||
}
|
}
|
||||||
|
dbMsgs[update.index].ReasoningContent = update.reasoningContent
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.InfoCF("seahorse", "bootstrap: repaired missing reasoning_content", map[string]any{
|
logger.InfoCF("seahorse", "bootstrap: repaired missing reasoning_content", map[string]any{
|
||||||
|
|
|
||||||
|
|
@ -759,6 +759,73 @@ func TestBootstrapRepairsMissingReasoningContentWithoutDroppingSummaries(t *test
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestBootstrapRepairsMissingReasoningContentOnPrefixBeforeAppendingDelta(t *testing.T) {
|
||||||
|
eng := newTestEngine(t)
|
||||||
|
ctx := context.Background()
|
||||||
|
sessionKey := "agent:repair-reasoning-prefix"
|
||||||
|
|
||||||
|
conv, err := eng.store.GetOrCreateConversation(ctx, sessionKey)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("GetOrCreateConversation: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
userMsg, err := eng.store.AddMessage(ctx, conv.ConversationID, "user", "hello", 3)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("AddMessage user: %v", err)
|
||||||
|
}
|
||||||
|
assistantMsg, err := eng.store.AddMessage(ctx, conv.ConversationID, "assistant", "world", 3)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("AddMessage assistant: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = eng.store.AppendContextMessages(
|
||||||
|
ctx,
|
||||||
|
conv.ConversationID,
|
||||||
|
[]int64{userMsg.ID, assistantMsg.ID},
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("AppendContextMessages: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
err = eng.Bootstrap(ctx, sessionKey, []Message{
|
||||||
|
{Role: "user", Content: "hello", TokenCount: 3},
|
||||||
|
{Role: "assistant", Content: "world", ReasoningContent: "let me think this through", TokenCount: 3},
|
||||||
|
{Role: "user", Content: "follow-up", TokenCount: 2},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Bootstrap: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
stored, err := eng.store.GetMessages(ctx, conv.ConversationID, 10, 0)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("GetMessages: %v", err)
|
||||||
|
}
|
||||||
|
if len(stored) != 3 {
|
||||||
|
t.Fatalf("stored messages = %d, want 3", len(stored))
|
||||||
|
}
|
||||||
|
if stored[1].ReasoningContent != "let me think this through" {
|
||||||
|
t.Errorf(
|
||||||
|
"stored[1].ReasoningContent = %q, want %q",
|
||||||
|
stored[1].ReasoningContent,
|
||||||
|
"let me think this through",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
if stored[2].Content != "follow-up" {
|
||||||
|
t.Errorf("stored[2].Content = %q, want %q", stored[2].Content, "follow-up")
|
||||||
|
}
|
||||||
|
|
||||||
|
items, err := eng.store.GetContextItems(ctx, conv.ConversationID)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("GetContextItems: %v", err)
|
||||||
|
}
|
||||||
|
if len(items) != 3 {
|
||||||
|
t.Fatalf("context items = %d, want 3", len(items))
|
||||||
|
}
|
||||||
|
if items[2].ItemType != "message" || items[2].MessageID != stored[2].ID {
|
||||||
|
t.Errorf("last context item = %+v, want appended message %d", items[2], stored[2].ID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestEngineBootstrapDelta(t *testing.T) {
|
func TestEngineBootstrapDelta(t *testing.T) {
|
||||||
eng := newTestEngine(t)
|
eng := newTestEngine(t)
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue