From 394c5494c9d78e04d31fa07395b8621d38f74674 Mon Sep 17 00:00:00 2001 From: Rahul Bansal Date: Sat, 21 Feb 2026 11:36:52 +0530 Subject: [PATCH] fix: doctor --fix now repairs corrupt sessions instead of deleting them Adds repairSessionMessages that injects synthetic tool_result messages for orphaned tool_use blocks and drops orphaned tool_result messages. Sessions are written back to disk with the repairs applied, preserving conversation history. --- pkg/doctor/doctor.go | 68 ++++++++++++++++++++++++++++++++++++--- pkg/doctor/doctor_test.go | 47 +++++++++++++++++++++++++++ 2 files changed, 111 insertions(+), 4 deletions(-) diff --git a/pkg/doctor/doctor.go b/pkg/doctor/doctor.go index a9b212dea..7dcdf508c 100644 --- a/pkg/doctor/doctor.go +++ b/pkg/doctor/doctor.go @@ -294,8 +294,8 @@ func checkSessions(opts Options) Result { for _, p := range problems { r.AddFixable(check, SeverityError, fmt.Sprintf("%s: %s", entry.Name(), p), - "remove corrupt session file", - makeSessionDeleteFunc(filePath), + "repair corrupt session (inject synthetic tool results, drop orphans)", + makeSessionRepairFunc(filePath, sess), ) } } @@ -380,9 +380,69 @@ func checkSessionMessages(msgs []providers.Message) []string { return problems } -func makeSessionDeleteFunc(path string) func() error { +// repairSessionMessages fixes orphaned tool_use/tool_result pairs in a message slice. +func repairSessionMessages(msgs []providers.Message) []providers.Message { + if len(msgs) == 0 { + return msgs + } + + // Collect all tool_call IDs + toolCallIDs := map[string]bool{} + for _, m := range msgs { + if m.Role == "assistant" { + for _, tc := range m.ToolCalls { + if tc.ID != "" { + toolCallIDs[tc.ID] = true + } + } + } + } + + // Collect all tool_result IDs + toolResultIDs := map[string]bool{} + for _, m := range msgs { + if m.Role == "tool" && m.ToolCallID != "" { + toolResultIDs[m.ToolCallID] = true + } + } + + // Build repaired slice: drop orphan results, inject missing results + repaired := make([]providers.Message, 0, len(msgs)) + for _, m := range msgs { + // Drop orphaned tool_result + if m.Role == "tool" && m.ToolCallID != "" && !toolCallIDs[m.ToolCallID] { + continue + } + repaired = append(repaired, m) + + // Inject missing tool_results after assistant+tool_calls + if m.Role == "assistant" && len(m.ToolCalls) > 0 { + for _, tc := range m.ToolCalls { + if tc.ID == "" { + continue + } + if !toolResultIDs[tc.ID] { + repaired = append(repaired, providers.Message{ + Role: "tool", + ToolCallID: tc.ID, + Content: "[tool result unavailable - session was repaired by picoclaw doctor]", + }) + toolResultIDs[tc.ID] = true + } + } + } + } + return repaired +} + +func makeSessionRepairFunc(path string, sess sessionFile) func() error { return func() error { - return os.Remove(path) + sess.Messages = repairSessionMessages(sess.Messages) + data, err := json.MarshalIndent(sess, "", " ") + if err != nil { + return fmt.Errorf("marshaling repaired session: %w", err) + } + return os.WriteFile(path, data, 0o644) } } diff --git a/pkg/doctor/doctor_test.go b/pkg/doctor/doctor_test.go index e2186d68f..42542ad34 100644 --- a/pkg/doctor/doctor_test.go +++ b/pkg/doctor/doctor_test.go @@ -88,6 +88,53 @@ func TestCheckSessionMessages_ConsecutiveUserMessages(t *testing.T) { } } +func TestRepairSessionMessages_InjectsResult(t *testing.T) { + msgs := []providers.Message{ + {Role: "user", Content: "hello"}, + { + Role: "assistant", Content: "checking", + ToolCalls: []providers.ToolCall{{ID: "call_1", Name: "exec"}}, + }, + {Role: "assistant", Content: "done"}, + } + repaired := repairSessionMessages(msgs) + if len(repaired) != 4 { + t.Fatalf("expected 4 messages, got %d", len(repaired)) + } + if repaired[2].Role != "tool" || repaired[2].ToolCallID != "call_1" { + t.Errorf("expected injected tool_result at index 2, got role=%q id=%q", + repaired[2].Role, repaired[2].ToolCallID) + } +} + +func TestRepairSessionMessages_DropsOrphanResult(t *testing.T) { + msgs := []providers.Message{ + {Role: "user", Content: "hello"}, + {Role: "tool", Content: "orphaned", ToolCallID: "call_x"}, + {Role: "assistant", Content: "hi"}, + } + repaired := repairSessionMessages(msgs) + if len(repaired) != 2 { + t.Fatalf("expected 2 messages, got %d", len(repaired)) + } +} + +func TestRepairSessionMessages_AlreadyClean(t *testing.T) { + msgs := []providers.Message{ + {Role: "user", Content: "hello"}, + { + Role: "assistant", Content: "checking", + ToolCalls: []providers.ToolCall{{ID: "call_1", Name: "exec"}}, + }, + {Role: "tool", Content: "output", ToolCallID: "call_1"}, + {Role: "assistant", Content: "done"}, + } + repaired := repairSessionMessages(msgs) + if len(repaired) != 4 { + t.Errorf("clean messages should be unchanged, got %d", len(repaired)) + } +} + func TestSeverityString(t *testing.T) { tests := []struct { sev Severity