diff --git a/internal/llmadapters/pi_rpc.go b/internal/llmadapters/pi_rpc.go index 36e1a32..75c36c2 100644 --- a/internal/llmadapters/pi_rpc.go +++ b/internal/llmadapters/pi_rpc.go @@ -2,6 +2,7 @@ package llmadapters import ( "bufio" + "bytes" "context" "encoding/json" "errors" @@ -743,16 +744,20 @@ func normalizePiRPCLogLine(line []byte) []byte { if !ok { return append(logLine, '\n') } - if compactPartial := compactPiRPCPartialForLog(partialRaw); len(compactPartial) > 0 { - assistantEvent["partial"] = compactPartial - } else { - delete(assistantEvent, "partial") + if !piRPCMessageUpdateRootIsRedundant(event, assistantEvent, partialRaw) { + return append(logLine, '\n') + } + compactPartial := compactPiRPCPartialForLog(partialRaw) + if len(compactPartial) == 0 { + return append(logLine, '\n') } + assistantEvent["partial"] = compactPartial normalizedAssistantEvent, err := json.Marshal(assistantEvent) if err != nil { return append(logLine, '\n') } event["assistantMessageEvent"] = normalizedAssistantEvent + delete(event, "message") normalized, err := json.Marshal(event) if err != nil { return append(logLine, '\n') @@ -760,6 +765,88 @@ func normalizePiRPCLogLine(line []byte) []byte { return append(normalized, '\n') } +func piRPCMessageUpdateRootIsRedundant(event, assistantEvent map[string]json.RawMessage, partialRaw json.RawMessage) bool { + if !validPiRPCProjection(assistantEvent, partialRaw) { + return false + } + rootRaw, ok := event["message"] + return ok && bytes.Equal(rootRaw, partialRaw) +} + +func validPiRPCProjection(assistantEvent map[string]json.RawMessage, partialRaw json.RawMessage) bool { + eventType := rawString(assistantEvent, "type") + if !isKnownPiRPCAssistantMessageEventType(eventType) || !piRPCJSONNumber(assistantEvent["contentIndex"]) { + return false + } + partial, ok := piRPCJSONObject(partialRaw) + if !ok || rawString(partial, "role") != "assistant" { + return false + } + for _, key := range []string{"role", "provider", "model", "api", "stopReason"} { + if !piRPCJSONString(partial[key]) { + return false + } + } + switch eventType { + case "text_delta", "thinking_delta", "toolcall_delta": + return piRPCJSONString(assistantEvent["delta"]) + case "text_end", "thinking_end": + return piRPCJSONString(assistantEvent["content"]) + case "toolcall_end": + return validPiRPCToolCall(assistantEvent["toolCall"]) + default: + return true + } +} + +func validPiRPCToolCall(value json.RawMessage) bool { + toolCall, ok := piRPCJSONObject(value) + if !ok || rawString(toolCall, "type") != "toolCall" { + return false + } + if !piRPCJSONString(toolCall["id"]) || !piRPCJSONString(toolCall["name"]) { + return false + } + _, ok = piRPCJSONObject(toolCall["arguments"]) + return ok +} + +func piRPCJSONObject(value json.RawMessage) (map[string]json.RawMessage, bool) { + if len(value) == 0 || bytes.Equal(bytes.TrimSpace(value), []byte("null")) { + return nil, false + } + var object map[string]json.RawMessage + if err := json.Unmarshal(value, &object); err != nil || object == nil { + return nil, false + } + return object, true +} + +func piRPCJSONString(value json.RawMessage) bool { + if len(value) == 0 || bytes.Equal(bytes.TrimSpace(value), []byte("null")) { + return false + } + var text string + return json.Unmarshal(value, &text) == nil +} + +func piRPCJSONNumber(value json.RawMessage) bool { + if len(value) == 0 || bytes.Equal(bytes.TrimSpace(value), []byte("null")) { + return false + } + var number float64 + return json.Unmarshal(value, &number) == nil +} + +func isKnownPiRPCAssistantMessageEventType(eventType string) bool { + switch eventType { + case "text_start", "text_delta", "text_end", "thinking_start", "thinking_delta", "thinking_end", "toolcall_start", "toolcall_delta", "toolcall_end": + return true + default: + return false + } +} + func compactPiRPCPartialForLog(partialRaw json.RawMessage) json.RawMessage { var partial map[string]json.RawMessage if err := json.Unmarshal(partialRaw, &partial); err != nil { diff --git a/internal/llmadapters/pi_rpc_test.go b/internal/llmadapters/pi_rpc_test.go index b5ebf2d..d8c6551 100644 --- a/internal/llmadapters/pi_rpc_test.go +++ b/internal/llmadapters/pi_rpc_test.go @@ -2,6 +2,7 @@ package llmadapters import ( "bufio" + "bytes" "context" "encoding/json" "errors" @@ -9,6 +10,7 @@ import ( "os" "os/exec" "path/filepath" + "reflect" "strconv" "strings" "testing" @@ -588,6 +590,9 @@ func TestPiRPCLogStripsCumulativeStreamingPartials(t *testing.T) { if string(response.StructuredOutput) != `{"ok":true}` { t.Fatalf("StructuredOutput = %s, want final assistant text", response.StructuredOutput) } + if stream.SessionID() != "session-1" { + t.Fatalf("SessionID = %q, want session-1", stream.SessionID()) + } logged, err := os.ReadFile(logPath) // #nosec G304 -- test reads the log path it created with t.TempDir. if err != nil { @@ -604,6 +609,7 @@ func TestPiRPCLogStripsCumulativeStreamingPartials(t *testing.T) { t.Fatalf("log = %s, want final RPC events preserved", logText) } assertPiRPCMessageUpdatesCompacted(t, logged) + assertPiRPCFinalEventsPreserved(t, logged) if len(logged) > 10_000 { t.Fatalf("log size = %d bytes, want normalized bounded log", len(logged)) } @@ -620,14 +626,39 @@ func assertPiRPCMessageUpdatesCompacted(t *testing.T, logged []byte) { if event["type"] != "message_update" { continue } + if _, ok := event["message"]; ok { + t.Fatalf("message_update = %#v, want redundant root message stripped", event) + } + if event["sessionId"] != "session-1" { + t.Fatalf("message_update sessionId = %#v, want session-1", event["sessionId"]) + } assistantEvent, ok := event["assistantMessageEvent"].(map[string]any) if !ok { t.Fatalf("message_update = %#v, want assistantMessageEvent", event) } + if assistantEvent["type"] != "thinking_delta" { + t.Fatalf("assistantMessageEvent.type = %#v, want thinking_delta", assistantEvent["type"]) + } + if assistantEvent["contentIndex"] != float64(0) { + t.Fatalf("assistantMessageEvent.contentIndex = %#v, want 0", assistantEvent["contentIndex"]) + } + if assistantEvent["delta"] != "x" { + t.Fatalf("assistantMessageEvent.delta = %#v, want x", assistantEvent["delta"]) + } partial, ok := assistantEvent["partial"].(map[string]any) if !ok { t.Fatalf("message_update = %#v, want compact partial metadata", event) } + wantPartial := map[string]any{ + "role": "assistant", + "provider": "opencode-go", + "model": "deepseek-v4-pro", + "api": "openai-completions", + "stopReason": "stop", + } + if !reflect.DeepEqual(partial, wantPartial) { + t.Fatalf("partial = %#v, want exact compact metadata %#v", partial, wantPartial) + } if _, ok := partial["content"]; ok { t.Fatalf("partial = %#v, want content stripped", partial) } @@ -640,6 +671,371 @@ func assertPiRPCMessageUpdatesCompacted(t *testing.T, logged []byte) { } } +func assertPiRPCFinalEventsPreserved(t *testing.T, logged []byte) { + t.Helper() + scanner := bufio.NewScanner(strings.NewReader(string(logged))) + var messageEndFound, agentEndFound bool + for scanner.Scan() { + var event map[string]any + if err := json.Unmarshal(scanner.Bytes(), &event); err != nil { + t.Fatalf("Unmarshal(log line): %v", err) + } + switch event["type"] { + case "message_end": + message, ok := event["message"].(map[string]any) + if !ok { + t.Fatalf("message_end = %#v, want message", event) + } + content, ok := message["content"].([]any) + if !ok || len(content) == 0 { + t.Fatalf("message_end.message = %#v, want final content", message) + } + contentBlock, ok := content[0].(map[string]any) + if !ok || contentBlock["text"] != `{"ok":true}` { + t.Fatalf("message_end.message = %#v, want final structured content", message) + } + usage, ok := message["usage"].(map[string]any) + if !ok || usage["input"] != float64(100) || usage["output"] != float64(50) { + t.Fatalf("message_end.message = %#v, want final usage", message) + } + messageEndFound = true + case "agent_end": + messages, ok := event["messages"].([]any) + if !ok || len(messages) == 0 { + t.Fatalf("agent_end = %#v, want messages", event) + } + agentEndFound = true + } + } + if err := scanner.Err(); err != nil { + t.Fatalf("Scan(log): %v", err) + } + if !messageEndFound { + t.Fatal("log missing message_end with final message content and usage") + } + if !agentEndFound { + t.Fatal("log missing agent_end messages") + } +} + +func TestNormalizePiRPCLogLineCompactsAllKnownPiAssistantEvents(t *testing.T) { + tests := []struct { + name string + eventType string + fields map[string]any + }{ + {name: "text start", eventType: "text_start"}, + {name: "text delta", eventType: "text_delta", fields: map[string]any{"delta": "text chunk"}}, + {name: "text end", eventType: "text_end", fields: map[string]any{"content": "complete text"}}, + {name: "thinking start", eventType: "thinking_start"}, + {name: "thinking delta", eventType: "thinking_delta", fields: map[string]any{"delta": "thinking chunk"}}, + {name: "thinking end", eventType: "thinking_end", fields: map[string]any{"content": "complete thinking"}}, + {name: "tool call start", eventType: "toolcall_start"}, + {name: "tool call delta", eventType: "toolcall_delta", fields: map[string]any{"delta": "{\"path\":\"x\"}"}}, + { + name: "tool call end", + eventType: "toolcall_end", + fields: map[string]any{"toolCall": map[string]any{ + "type": "toolCall", "id": "call-1", "name": "Read", "arguments": map[string]any{"path": "x"}, + }}, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + event, _, assistantEvent := piRPCProjectionEventFixture(tt.eventType) + for key, value := range tt.fields { + assistantEvent[key] = value + } + line, err := json.Marshal(event) + if err != nil { + t.Fatalf("Marshal(event): %v", err) + } + normalized := normalizePiRPCLogLine(line) + gotAssistantEvent := assertPiRPCProjectionCompacted(t, normalized, tt.eventType, piRPCProjectionPartialFixture()) + for key, want := range tt.fields { + if !reflect.DeepEqual(gotAssistantEvent[key], want) { + t.Fatalf("assistantMessageEvent[%q] = %#v, want %#v", key, gotAssistantEvent[key], want) + } + } + }) + } +} + +func TestNormalizePiRPCLogLinePreservesOriginalForProjectionFailures(t *testing.T) { + tests := []struct { + name string + mutate func(event, partial, assistantEvent map[string]any) + }{ + {name: "missing role", mutate: piRPCDeleteProjectionField("role")}, + {name: "missing provider", mutate: piRPCDeleteProjectionField("provider")}, + {name: "missing model", mutate: piRPCDeleteProjectionField("model")}, + {name: "missing api", mutate: piRPCDeleteProjectionField("api")}, + {name: "missing stop reason", mutate: piRPCDeleteProjectionField("stopReason")}, + {name: "non-string role", mutate: piRPCSetProjectionField("role", 1)}, + {name: "non-string provider", mutate: piRPCSetProjectionField("provider", 1)}, + {name: "non-string model", mutate: piRPCSetProjectionField("model", 1)}, + {name: "non-string api", mutate: piRPCSetProjectionField("api", 1)}, + {name: "non-string stop reason", mutate: piRPCSetProjectionField("stopReason", 1)}, + { + name: "invalid role", + mutate: func(_, partial, _ map[string]any) { partial["role"] = "user" }, + }, + { + name: "unknown event type", + mutate: func(_, _, assistantEvent map[string]any) { assistantEvent["type"] = "future_event" }, + }, + { + name: "missing content index", + mutate: func(_, _, assistantEvent map[string]any) { delete(assistantEvent, "contentIndex") }, + }, + { + name: "non-numeric content index", + mutate: func(_, _, assistantEvent map[string]any) { assistantEvent["contentIndex"] = "0" }, + }, + { + name: "missing delta", + mutate: func(_, _, assistantEvent map[string]any) { + assistantEvent["type"] = "text_delta" + delete(assistantEvent, "delta") + }, + }, + { + name: "non-string delta", + mutate: func(_, _, assistantEvent map[string]any) { + assistantEvent["type"] = "text_delta" + assistantEvent["delta"] = 1 + }, + }, + { + name: "missing end content", + mutate: func(_, _, assistantEvent map[string]any) { + assistantEvent["type"] = "text_end" + delete(assistantEvent, "content") + }, + }, + { + name: "non-string end content", + mutate: func(_, _, assistantEvent map[string]any) { + assistantEvent["type"] = "thinking_end" + assistantEvent["content"] = 1 + }, + }, + { + name: "missing tool call", + mutate: func(_, _, assistantEvent map[string]any) { + assistantEvent["type"] = "toolcall_end" + delete(assistantEvent, "toolCall") + }, + }, + { + name: "invalid tool call", + mutate: func(_, _, assistantEvent map[string]any) { + assistantEvent["type"] = "toolcall_end" + assistantEvent["toolCall"] = map[string]any{"type": "toolCall", "id": "call-1"} + }, + }, + { + name: "missing partial", + mutate: func(_, _, assistantEvent map[string]any) { delete(assistantEvent, "partial") }, + }, + { + name: "non-object partial", + mutate: func(_, _, assistantEvent map[string]any) { assistantEvent["partial"] = []any{} }, + }, + { + name: "partial without projection keys", + mutate: func(_, partial, _ map[string]any) { + for _, key := range []string{"role", "provider", "model", "api", "stopReason"} { + delete(partial, key) + } + }, + }, + { + name: "missing root message", + mutate: func(event, _, _ map[string]any) { delete(event, "message") }, + }, + { + name: "nonidentical root message", + mutate: func(event, partial, _ map[string]any) { + root := make(map[string]any, len(partial)) + for key, value := range partial { + root[key] = value + } + root["stopReason"] = "different" + event["message"] = root + }, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + event, partial, assistantEvent := piRPCProjectionEventFixture("thinking_delta") + tt.mutate(event, partial, assistantEvent) + line, err := json.Marshal(event) + if err != nil { + t.Fatalf("Marshal(event): %v", err) + } + normalized := normalizePiRPCLogLine(line) + want := append(append([]byte(nil), line...), '\n') + if !bytes.Equal(normalized, want) { + t.Fatalf("normalized = %s, want original line preserved", normalized) + } + }) + } +} + +func TestNormalizePiRPCLogLinePreservesOriginalForDifferentlyEncodedEqualRoot(t *testing.T) { + line := []byte(`{"type":"message_update","message":{"role": "assistant", "provider": "opencode-go", "model": "deepseek-v4-pro", "api": "openai-completions", "stopReason": "stop", "content": [], "usage": {"input": 100}, "timestamp": 1700000000000},"assistantMessageEvent":{"type":"thinking_delta","contentIndex":0,"delta":"x","partial":{"timestamp":1700000000000,"usage":{"input":100},"content":[],"stopReason":"stop","api":"openai-completions","model":"deepseek-v4-pro","provider":"opencode-go","role":"assistant"}}}`) + + var event map[string]json.RawMessage + if err := json.Unmarshal(line, &event); err != nil { + t.Fatalf("Unmarshal(line): %v", err) + } + var root, partial any + if err := json.Unmarshal(event["message"], &root); err != nil { + t.Fatalf("Unmarshal(root): %v", err) + } + var assistantEvent map[string]json.RawMessage + if err := json.Unmarshal(event["assistantMessageEvent"], &assistantEvent); err != nil { + t.Fatalf("Unmarshal(assistantMessageEvent): %v", err) + } + if err := json.Unmarshal(assistantEvent["partial"], &partial); err != nil { + t.Fatalf("Unmarshal(partial): %v", err) + } + if !reflect.DeepEqual(root, partial) { + t.Fatalf("root and partial differ semantically: root=%#v partial=%#v", root, partial) + } + + normalized := normalizePiRPCLogLine(line) + want := append(append([]byte(nil), line...), '\n') + if !bytes.Equal(normalized, want) { + t.Fatalf("normalized = %s, want exact original line preserved", normalized) + } +} + +func TestNormalizePiRPCLogLineCompactsProjectionWithMalformedCumulativeFields(t *testing.T) { + event, partial, _ := piRPCProjectionEventFixture("thinking_delta") + partial["content"] = map[string]any{"unexpected": true} + partial["usage"] = []any{"unexpected"} + partial["timestamp"] = "not-a-number" + line, err := json.Marshal(event) + if err != nil { + t.Fatalf("Marshal(event): %v", err) + } + assertPiRPCProjectionCompacted(t, normalizePiRPCLogLine(line), "thinking_delta", piRPCProjectionPartialFixture()) +} + +func TestNormalizePiRPCLogLineCompactsStringProjectionFieldsWithoutPolicyChecks(t *testing.T) { + event, partial, _ := piRPCProjectionEventFixture("thinking_delta") + partial["provider"] = "" + partial["model"] = "" + partial["api"] = "" + partial["stopReason"] = "future-stop-reason" + line, err := json.Marshal(event) + if err != nil { + t.Fatalf("Marshal(event): %v", err) + } + assertPiRPCProjectionCompacted(t, normalizePiRPCLogLine(line), "thinking_delta", map[string]any{ + "role": "assistant", "provider": "", "model": "", "api": "", "stopReason": "future-stop-reason", + }) +} + +func assertPiRPCProjectionCompacted(t *testing.T, normalized []byte, eventType string, wantPartial map[string]any) map[string]any { + t.Helper() + var event map[string]any + if err := json.Unmarshal(normalized, &event); err != nil { + t.Fatalf("Unmarshal(normalized): %v", err) + } + if _, ok := event["message"]; ok { + t.Fatalf("normalized message_update = %#v, want redundant root message stripped", event) + } + if event["sessionId"] != "session-1" || event["preservedSibling"] != "keep" { + t.Fatalf("normalized event siblings = %#v, want session and sibling preserved", event) + } + assistantEvent, ok := event["assistantMessageEvent"].(map[string]any) + if !ok { + t.Fatalf("normalized message_update = %#v, want assistantMessageEvent", event) + } + if assistantEvent["type"] != eventType || assistantEvent["contentIndex"] != float64(0) || assistantEvent["preservedSibling"] != "keep" { + t.Fatalf("assistantMessageEvent = %#v, want retained event fields", assistantEvent) + } + partial, ok := assistantEvent["partial"].(map[string]any) + if !ok { + t.Fatalf("assistantMessageEvent = %#v, want compact partial", assistantEvent) + } + if !reflect.DeepEqual(partial, wantPartial) { + t.Fatalf("partial = %#v, want exact compact projection %#v", partial, wantPartial) + } + return assistantEvent +} + +func piRPCProjectionPartialFixture() map[string]any { + return map[string]any{ + "role": "assistant", + "provider": "opencode-go", + "model": "deepseek-v4-pro", + "api": "openai-completions", + "stopReason": "stop", + } +} + +func piRPCProjectionEventFixture(eventType string) (map[string]any, map[string]any, map[string]any) { + partial := map[string]any{ + "role": "assistant", + "provider": "opencode-go", + "model": "deepseek-v4-pro", + "api": "openai-completions", + "stopReason": "stop", + "timestamp": 1700000000000, + "content": []map[string]any{{"type": "thinking", "thinking": "cumulative"}}, + "usage": map[string]any{ + "input": 100, + "output": 50, + "cacheRead": 0, + "cacheWrite": 0, + "totalTokens": 150, + "cost": map[string]any{ + "input": 0, + "output": 0, + "cacheRead": 0, + "cacheWrite": 0, + "total": 0, + }, + }, + } + assistantEvent := map[string]any{ + "type": eventType, + "contentIndex": 0, + "partial": partial, + "preservedSibling": "keep", + } + switch eventType { + case "text_delta", "thinking_delta", "toolcall_delta": + assistantEvent["delta"] = "x" + case "text_end", "thinking_end": + assistantEvent["content"] = "complete" + case "toolcall_end": + assistantEvent["toolCall"] = map[string]any{ + "type": "toolCall", "id": "call-1", "name": "Read", "arguments": map[string]any{"path": "x"}, + } + } + event := map[string]any{ + "type": "message_update", + "message": partial, + "sessionId": "session-1", + "preservedSibling": "keep", + "assistantMessageEvent": assistantEvent, + } + return event, partial, assistantEvent +} + +func piRPCDeleteProjectionField(key string) func(event, partial, assistantEvent map[string]any) { + return func(_, partial, _ map[string]any) { delete(partial, key) } +} + +func piRPCSetProjectionField(key string, value any) func(event, partial, assistantEvent map[string]any) { + return func(_, partial, _ map[string]any) { partial[key] = value } +} + func TestPiRPCProtocolFailures(t *testing.T) { t.Run("prompt response failure", func(t *testing.T) { recordPath := filepath.Join(t.TempDir(), "record.json") @@ -938,26 +1334,44 @@ func TestPiRPCHelperProcess(_ *testing.T) { thinking := "" for i := 0; i < 25; i++ { thinking += strings.Repeat("x", 200) + partial := map[string]any{ + "role": "assistant", + "provider": "opencode-go", + "model": "deepseek-v4-pro", + "api": "openai-completions", + "stopReason": "stop", + "timestamp": 1700000000000, + "content": []map[string]any{{"type": "thinking", "thinking": thinking}}, + "usage": map[string]any{ + "input": 100, + "output": 50, + "cacheRead": 0, + "cacheWrite": 0, + "totalTokens": 150, + "cost": map[string]any{ + "input": 0, + "output": 0, + "cacheRead": 0, + "cacheWrite": 0, + "total": 0, + }, + }, + } event := map[string]any{ - "type": "message_update", + "type": "message_update", + "message": partial, + "sessionId": "session-1", "assistantMessageEvent": map[string]any{ "type": "thinking_delta", "contentIndex": 0, "delta": "x", - "partial": map[string]any{ - "role": "assistant", - "provider": "opencode-go", - "model": "deepseek-v4-pro", - "stopReason": "stop", - "content": []map[string]any{{"type": "thinking", "thinking": thinking}}, - "usage": map[string]any{"tokensIn": 100, "tokensOut": 50, "totalTokens": 150}, - }, + "partial": partial, }, } data, _ := json.Marshal(event) fmt.Println(string(data)) } - fmt.Println(`{"type":"message_end","message":{"role":"assistant","content":[{"type":"text","text":"{\"ok\":true}"}],"usage":{"tokensIn":100,"tokensOut":50}}}`) + fmt.Println(`{"type":"message_end","message":{"role":"assistant","content":[{"type":"text","text":"{\"ok\":true}"}],"api":"openai-completions","provider":"opencode-go","model":"deepseek-v4-pro","usage":{"input":100,"output":50,"cacheRead":0,"cacheWrite":0,"totalTokens":150,"cost":{"input":0,"output":0,"cacheRead":0,"cacheWrite":0,"total":0}},"stopReason":"stop","timestamp":1700000000000}}`) fmt.Println(`{"type":"agent_end","messages":[{"role":"assistant","content":[{"type":"text","text":"{\"ok\":true}"}]}]}`) case "prompt-failure": fmt.Println(`{"id":"prompt-1","type":"response","command":"prompt","success":false,"error":"No API key found for opencode-go"}`)