From ecdf0c44d139a4826ec30051e137812e176e79a0 Mon Sep 17 00:00:00 2001 From: m1saka Date: Sun, 23 Aug 2026 22:16:32 +0800 Subject: [PATCH] Fix upstream 400: flatten multi-part content arrays and sanitize invalid tool_calls MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The Zhanlu upstream gateway returns HTTP 400 (请求消息格式错误) for two message-structure issues that AI coding tools like MiMoCode produce: 1. Multi-part content arrays — OpenAI SDKs send content as [{type:"text",text:"…"}] arrays; the gateway only accepts string content. normalizeMessages now flattens text-only arrays into a concatenated string (non-text parts like images are preserved). 2. Invalid tool_calls — when a tool call fails, MiMoCode emits {name:"invalid", arguments:{"tool":"task","error":"…"}} placeholders. The gateway rejects function names not in the tools list and non-object arguments (e.g. "-1"). normalizeMessages now sanitizes these in place: the real tool name is extracted from the arguments' "tool" field (falling back to an arbitrary declared tool), and non-JSON-object arguments are replaced with "{}". No messages or tool results are removed, preserving the full conversation context including error feedback. Verified with the exact error.md request: 0/10 400 errors after fix (vs 10/10 before). Model returns valid streaming responses with task tool calls. --- .gitignore | 1 + internal/openai/types.go | 183 ++++++++++++- internal/openai/types_test.go | 486 ++++++++++++++++++++++++++++++++++ 3 files changed, 669 insertions(+), 1 deletion(-) create mode 100644 internal/openai/types_test.go diff --git a/.gitignore b/.gitignore index 1ce7795..074c5cb 100644 --- a/.gitignore +++ b/.gitignore @@ -14,3 +14,4 @@ tmp/ temp/ source/ output/ +error.md diff --git a/internal/openai/types.go b/internal/openai/types.go index 85c5ea5..28b65bb 100644 --- a/internal/openai/types.go +++ b/internal/openai/types.go @@ -33,7 +33,7 @@ func (r *ChatCompletionRequest) UnmarshalJSON(data []byte) error { func (r ChatCompletionRequest) MarshalForUpstream() ([]byte, error) { m := map[string]any{ "model": r.Model, - "messages": r.Messages, + "messages": normalizeMessages(r.Messages, extractToolNames(r.Extra)), "stream": r.Stream, "stream_options": map[string]any{"include_usage": true}, } @@ -51,6 +51,187 @@ func (r ChatCompletionRequest) MarshalForUpstream() ([]byte, error) { return json.Marshal(m) } +// extractToolNames returns the set of function names declared in the tools +// array (Extra["tools"]), used to validate tool_calls in the message history. +func extractToolNames(extra map[string]json.RawMessage) map[string]bool { + names := map[string]bool{} + raw, ok := extra["tools"] + if !ok { + return names + } + var tools []any + if json.Unmarshal(raw, &tools) != nil { + return names + } + for _, t := range tools { + tm, ok := t.(map[string]any) + if !ok { + continue + } + // Chat Completions format: {type:"function",function:{name:"…"}} + if fn, ok := tm["function"].(map[string]any); ok { + if name, _ := fn["name"].(string); name != "" { + names[name] = true + } + } + // Responses API flat format: {type:"function",name:"…"} + if name, _ := tm["name"].(string); name != "" { + names[name] = true + } + } + return names +} + +// normalizeMessages prepares the message history for the Zhanlu upstream +// gateway. It performs two normalizations: +// +// 1. Flattens OpenAI multi-part content — an array of {type:"text",text:"…"} +// parts — into a plain string. The gateway only accepts string content +// and returns HTTP 400 for arrays. +// +// 2. Sanitizes tool_calls that the gateway would reject: calls whose +// function name is not in the declared tools set (e.g. "invalid" +// placeholders from AI coding tools when a tool call fails), or whose +// arguments are not a valid JSON object. Instead of removing these +// calls (which would lose error feedback the model needs), the function +// name is replaced with a valid one and the arguments are replaced with +// "{}". The corresponding tool result messages are preserved so the +// model still sees the full conversation including errors. +// +// Both []any (from JSON decoding of /v1/chat/completions) and +// []map[string]any (from convertContentParts in the Responses API path) +// are handled for content arrays. +func normalizeMessages(messages []map[string]any, validToolNames map[string]bool) []map[string]any { + fallbackName := pickFallbackToolName(validToolNames) + + for _, msg := range messages { + // --- flatten multi-part text content --- + parts := contentAsAnySlice(msg["content"]) + if len(parts) > 0 { + var sb strings.Builder + allText := true + for _, p := range parts { + part, ok := p.(map[string]any) + if !ok { + allText = false + break + } + pt, _ := part["type"].(string) + if pt != "text" && pt != "input_text" && pt != "output_text" { + allText = false + break + } + text, _ := part["text"].(string) + sb.WriteString(text) + } + if allText { + msg["content"] = sb.String() + } + } + + // --- sanitize tool_calls in place --- + tcs := contentAsAnySlice(msg["tool_calls"]) + for _, tc := range tcs { + if tcMap, ok := tc.(map[string]any); ok { + sanitizeToolCall(tcMap, validToolNames, fallbackName) + } + } + } + return messages +} + +// sanitizeToolCall fixes a tool_call in place so the upstream gateway +// accepts it. Two fields are corrected: +// +// - function.name: if the name is not among the declared tools (when the +// set is non-empty), it is replaced. The replacement is extracted from +// the arguments' "tool" field (MiMoCode stores the real tool name there +// in error placeholders); failing that, an arbitrary declared tool name +// is used. +// +// - function.arguments: if the value is not a valid JSON object string +// (e.g. "-1", "", "true"), it is replaced with "{}". +// +// The tool_call id, type, and the tool result messages are left untouched, +// preserving the full conversation context for the model. +func sanitizeToolCall(tc map[string]any, validToolNames map[string]bool, fallbackName string) { + fn, ok := tc["function"].(map[string]any) + if !ok { + return + } + // Fix function name if not in the valid set. + name, _ := fn["name"].(string) + if len(validToolNames) > 0 && !validToolNames[name] { + args, _ := fn["arguments"].(string) + if extracted := extractToolNameFromArgs(args); validToolNames[extracted] { + fn["name"] = extracted + } else if fallbackName != "" { + fn["name"] = fallbackName + } + } + // Fix arguments if not a valid JSON object. + args, _ := fn["arguments"].(string) + if !isValidJSONObject(args) { + fn["arguments"] = "{}" + } +} + +// extractToolNameFromArgs attempts to read a "tool" field from the JSON +// object in args. MiMoCode's error placeholders store the real tool name +// here (e.g. {"tool":"task","error":"…"}). +func extractToolNameFromArgs(args string) string { + var m map[string]any + if json.Unmarshal([]byte(args), &m) != nil { + return "" + } + tool, _ := m["tool"].(string) + return tool +} + +// pickFallbackToolName returns an arbitrary name from the set for use as a +// last-resort replacement when no real tool name can be extracted. +func pickFallbackToolName(names map[string]bool) string { + for name := range names { + return name + } + return "" +} + +// isValidJSONObject reports whether s is a non-empty JSON string that decodes +// to a JSON object (map[string]any). The OpenAI tool_call arguments field +// must be a JSON object string like {"key":"value"}; the Zhanlu gateway +// rejects values like "-1", "true", or "" with HTTP 400. +func isValidJSONObject(s string) bool { + s = strings.TrimSpace(s) + if s == "" { + return false + } + var v any + if err := json.Unmarshal([]byte(s), &v); err != nil { + return false + } + _, ok := v.(map[string]any) + return ok +} + +// contentAsAnySlice returns the content value as a []any, handling both +// []any (JSON decoding) and []map[string]any (convertContentParts). It +// returns nil for non-slice values (strings, nil, etc.). +func contentAsAnySlice(v any) []any { + switch s := v.(type) { + case []any: + return s + case []map[string]any: + result := make([]any, len(s)) + for i, m := range s { + result[i] = m + } + return result + default: + return nil + } +} + type ErrorResponse struct { Error ErrorBody `json:"error"` } diff --git a/internal/openai/types_test.go b/internal/openai/types_test.go new file mode 100644 index 0000000..057b106 --- /dev/null +++ b/internal/openai/types_test.go @@ -0,0 +1,486 @@ +package openai + +import ( + "encoding/json" + "strings" + "testing" +) + +func TestMarshalForUpstream_StringContent(t *testing.T) { + req := ChatCompletionRequest{ + Model: "GLM-4.7", + Messages: []map[string]any{{"role": "user", "content": "hello"}}, + Stream: true, + } + body, err := req.MarshalForUpstream() + if err != nil { + t.Fatal(err) + } + var m map[string]any + if err := json.Unmarshal(body, &m); err != nil { + t.Fatal(err) + } + messages, _ := m["messages"].([]any) + msg, _ := messages[0].(map[string]any) + if content, _ := msg["content"].(string); content != "hello" { + t.Fatalf("content = %v, want %q", msg["content"], "hello") + } +} + +// TestMarshalForUpstream_FlattensArrayContent verifies that multi-part +// content arrays (the format used by OpenAI SDKs and AI coding tools) are +// flattened into a plain string so the Zhanlu upstream gateway accepts the +// request instead of returning HTTP 400. +func TestMarshalForUpstream_FlattensArrayContent(t *testing.T) { + req := ChatCompletionRequest{ + Model: "zhanlu/deepseek-v4-pro", + Messages: []map[string]any{ + {"role": "user", "content": []any{ + map[string]any{"type": "text", "text": "first part "}, + map[string]any{"type": "text", "text": "second part"}, + }}, + }, + Stream: true, + } + body, err := req.MarshalForUpstream() + if err != nil { + t.Fatal(err) + } + var m map[string]any + if err := json.Unmarshal(body, &m); err != nil { + t.Fatal(err) + } + messages, _ := m["messages"].([]any) + msg, _ := messages[0].(map[string]any) + content, ok := msg["content"].(string) + if !ok { + t.Fatalf("content type = %T, want string", msg["content"]) + } + if content != "first part second part" { + t.Fatalf("content = %q, want %q", content, "first part second part") + } +} + +// TestMarshalForUpstream_PreservesNonTextArrayContent verifies that content +// arrays containing non-text parts (e.g. images) are left intact rather than +// being flattened, so the upstream can handle multimodal content. +func TestMarshalForUpstream_PreservesNonTextArrayContent(t *testing.T) { + req := ChatCompletionRequest{ + Model: "GLM-4.7", + Messages: []map[string]any{ + {"role": "user", "content": []any{ + map[string]any{"type": "text", "text": "describe this"}, + map[string]any{"type": "image_url", "image_url": map[string]any{"url": "data:image/png;base64,abc"}}, + }}, + }, + Stream: true, + } + body, err := req.MarshalForUpstream() + if err != nil { + t.Fatal(err) + } + var m map[string]any + if err := json.Unmarshal(body, &m); err != nil { + t.Fatal(err) + } + messages, _ := m["messages"].([]any) + msg, _ := messages[0].(map[string]any) + parts, ok := msg["content"].([]any) + if !ok { + t.Fatalf("content type = %T, want []any (preserved array)", msg["content"]) + } + if len(parts) != 2 { + t.Fatalf("parts = %d, want 2", len(parts)) + } +} + +// TestMarshalForUpstream_ExtraFields verifies that extra fields (max_tokens, +// tools, etc.) are passed through to the upstream body. +func TestMarshalForUpstream_ExtraFields(t *testing.T) { + req := ChatCompletionRequest{ + Model: "GLM-4.7", + Stream: true, + Extra: map[string]json.RawMessage{ + "max_tokens": json.RawMessage(`32000`), + "tool_choice": json.RawMessage(`"auto"`), + }, + } + body, err := req.MarshalForUpstream() + if err != nil { + t.Fatal(err) + } + var m map[string]any + if err := json.Unmarshal(body, &m); err != nil { + t.Fatal(err) + } + if v, _ := m["max_tokens"].(float64); v != 32000 { + t.Fatalf("max_tokens = %v, want 32000", m["max_tokens"]) + } + if v, _ := m["tool_choice"].(string); v != "auto" { + t.Fatalf("tool_choice = %v, want \"auto\"", m["tool_choice"]) + } +} + +// TestMarshalForUpstream_GLMToolStream verifies that tool_stream is added +// for GLM models but not for non-GLM models. +func TestMarshalForUpstream_GLMToolStream(t *testing.T) { + cases := []struct { + model string + wantTool bool + }{ + {"GLM-4.7", true}, + {"zhanlu/glm-4-pro", true}, + {"zhanlu/deepseek-v4-pro", false}, + {"deepseek-chat", false}, + } + for _, tc := range cases { + req := ChatCompletionRequest{Model: tc.model, Stream: true} + body, err := req.MarshalForUpstream() + if err != nil { + t.Fatal(err) + } + var m map[string]any + _ = json.Unmarshal(body, &m) + _, has := m["tool_stream"] + if has != tc.wantTool { + t.Fatalf("model %q: tool_stream present = %v, want %v", tc.model, has, tc.wantTool) + } + } +} + +// TestNormalizeMessages_MixedContent verifies a mix of string and array +// content across multiple messages in a single request. +func TestNormalizeMessages_MixedContent(t *testing.T) { + messages := []map[string]any{ + {"role": "system", "content": "you are helpful"}, + {"role": "user", "content": []any{ + map[string]any{"type": "text", "text": "part 1 "}, + map[string]any{"type": "text", "text": "part 2"}, + }}, + {"role": "assistant", "content": "ok"}, + {"role": "user", "content": []any{ + map[string]any{"type": "text", "text": "only part"}, + }}, + } + result := normalizeMessages(messages, nil) + // system message unchanged + if c, _ := result[0]["content"].(string); c != "you are helpful" { + t.Fatalf("msg[0] content = %v", result[0]["content"]) + } + // array content flattened + if c, _ := result[1]["content"].(string); c != "part 1 part 2" { + t.Fatalf("msg[1] content = %v, want %q", result[1]["content"], "part 1 part 2") + } + // string content unchanged + if c, _ := result[2]["content"].(string); c != "ok" { + t.Fatalf("msg[2] content = %v", result[2]["content"]) + } + // single-part array flattened to string + if c, _ := result[3]["content"].(string); c != "only part" { + t.Fatalf("msg[3] content = %v, want %q", result[3]["content"], "only part") + } +} + +// TestNormalizeMessages_EmptyArray verifies that an empty content array +// becomes an empty string. +func TestNormalizeMessages_EmptyArray(t *testing.T) { + messages := []map[string]any{ + {"role": "user", "content": []any{}}, + } + result := normalizeMessages(messages, nil) + if c, _ := result[0]["content"].(string); c != "" { + t.Fatalf("content = %v, want empty string", result[0]["content"]) + } +} + +// TestNormalizeMessages_MapSliceContent verifies that content stored as +// []map[string]any (produced by convertContentParts in the Responses API +// path) is also flattened, not just []any (from JSON decoding). +func TestNormalizeMessages_MapSliceContent(t *testing.T) { + messages := []map[string]any{ + {"role": "user", "content": []map[string]any{ + {"type": "text", "text": "map part 1 "}, + {"type": "text", "text": "map part 2"}, + }}, + } + result := normalizeMessages(messages, nil) + if c, _ := result[0]["content"].(string); c != "map part 1 map part 2" { + t.Fatalf("content = %v, want %q", result[0]["content"], "map part 1 map part 2") + } +} + +// TestMarshalForUpstream_ErrorRequestReplay is a regression test mirroring +// the real failing request from error.md: a multi-turn conversation with +// system prompt, user messages with multi-part content arrays, assistant +// messages with tool_calls and reasoning_content, and tool result messages. +// It verifies that no message has array content after marshaling. +func TestMarshalForUpstream_ErrorRequestReplay(t *testing.T) { + req := ChatCompletionRequest{ + Model: "zhanlu/deepseek-v4-pro", + Messages: []map[string]any{ + {"role": "system", "content": strings.Repeat("system prompt ", 100)}, + {"role": "user", "content": []any{ + map[string]any{"type": "text", "text": "user question"}, + map[string]any{"type": "text", "text": "skill search"}, + }}, + {"role": "assistant", "content": "I will search.", "reasoning_content": "thinking...", "tool_calls": []any{ + map[string]any{"id": "call_1", "type": "function", "function": map[string]any{"name": "bash", "arguments": `{"command":"ls"}`}}, + }}, + {"role": "tool", "tool_call_id": "call_1", "content": "file1\nfile2"}, + {"role": "user", "content": []any{ + map[string]any{"type": "text", "text": "continue"}, + }}, + }, + Stream: true, + Extra: map[string]json.RawMessage{ + "max_tokens": json.RawMessage(`32000`), + "tool_choice": json.RawMessage(`"auto"`), + }, + } + body, err := req.MarshalForUpstream() + if err != nil { + t.Fatal(err) + } + var m map[string]any + if err := json.Unmarshal(body, &m); err != nil { + t.Fatal(err) + } + messages, _ := m["messages"].([]any) + for i, raw := range messages { + msg, _ := raw.(map[string]any) + switch c := msg["content"].(type) { + case string: + // OK — flattened + case []any: + t.Fatalf("message [%d] still has array content after marshal", i) + default: + t.Fatalf("message [%d] content type = %T", i, c) + } + } +} + +// TestNormalizeMessages_SanitizesInvalidToolCallNames verifies that +// tool_calls whose function name is not in the declared tools set are +// fixed in place — the name is replaced with a valid one — rather than +// removed. This preserves the conversation context including error +// feedback. MiMoCode emits {name:"invalid"} placeholders when a tool +// call fails; the Zhanlu gateway rejects unknown function names with +// HTTP 400. +func TestNormalizeMessages_SanitizesInvalidToolCallNames(t *testing.T) { + validNames := map[string]bool{"bash": true, "read": true} + messages := []map[string]any{ + {"role": "user", "content": "hello"}, + {"role": "assistant", "content": "I'll run bash.", "tool_calls": []any{ + map[string]any{"id": "call_ok", "type": "function", "function": map[string]any{"name": "bash", "arguments": `{"command":"ls"}`}}, + // "invalid" name with "tool" field in args — should be extracted + map[string]any{"id": "call_bad", "type": "function", "function": map[string]any{"name": "invalid", "arguments": `{"tool":"read","error":"failed"}`}}, + }}, + {"role": "tool", "tool_call_id": "call_ok", "content": "file1"}, + {"role": "tool", "tool_call_id": "call_bad", "content": "error result"}, + {"role": "user", "content": "thanks"}, + } + result := normalizeMessages(messages, validNames) + // All messages preserved (no removal) + if len(result) != 5 { + t.Fatalf("messages = %d, want 5 (sanitize preserves all)", len(result)) + } + // Both tool_calls kept; the invalid one's name should be fixed to "read" + // (extracted from args.tool) + asst := result[1] + tcs, _ := asst["tool_calls"].([]any) + if len(tcs) != 2 { + t.Fatalf("tool_calls = %d items, want 2 (preserved)", len(tcs)) + } + tc1, _ := tcs[0].(map[string]any) + fn1, _ := tc1["function"].(map[string]any) + if name, _ := fn1["name"].(string); name != "bash" { + t.Fatalf("first tool_call name = %q, want %q", name, "bash") + } + tc2, _ := tcs[1].(map[string]any) + fn2, _ := tc2["function"].(map[string]any) + if name, _ := fn2["name"].(string); name != "read" { + t.Fatalf("sanitized tool_call name = %q, want %q", name, "read") + } + // Tool results preserved + toolMsg := result[3] + if id, _ := toolMsg["tool_call_id"].(string); id != "call_bad" { + t.Fatalf("tool_call_id = %q, want %q", id, "call_bad") + } +} + +// TestNormalizeMessages_SanitizesInvalidArguments verifies that tool_calls +// whose arguments are not a valid JSON object string (e.g. "-1", "", +// "true") are fixed to "{}" in place, not removed. The Zhanlu gateway +// requires arguments to be a JSON object. +func TestNormalizeMessages_SanitizesInvalidArguments(t *testing.T) { + validNames := map[string]bool{"task": true} + messages := []map[string]any{ + {"role": "user", "content": "create tasks"}, + {"role": "assistant", "content": "creating.", "tool_calls": []any{ + map[string]any{"id": "call_good", "type": "function", "function": map[string]any{"name": "task", "arguments": `{"operation":"create","summary":"test"}`}}, + map[string]any{"id": "call_bad1", "type": "function", "function": map[string]any{"name": "task", "arguments": "-1"}}, + map[string]any{"id": "call_bad2", "type": "function", "function": map[string]any{"name": "task", "arguments": ""}}, + }}, + {"role": "tool", "tool_call_id": "call_good", "content": "created"}, + {"role": "tool", "tool_call_id": "call_bad1", "content": "error"}, + {"role": "tool", "tool_call_id": "call_bad2", "content": "error"}, + } + result := normalizeMessages(messages, validNames) + // All 5 messages preserved + if len(result) != 5 { + t.Fatalf("messages = %d, want 5 (sanitize preserves all)", len(result)) + } + asst := result[1] + tcs, _ := asst["tool_calls"].([]any) + if len(tcs) != 3 { + t.Fatalf("tool_calls = %d items, want 3 (preserved)", len(tcs)) + } + // Good args unchanged + tc0, _ := tcs[0].(map[string]any) + fn0, _ := tc0["function"].(map[string]any) + if args, _ := fn0["arguments"].(string); args != `{"operation":"create","summary":"test"}` { + t.Fatalf("good args changed: %q", args) + } + // Bad args fixed to "{}" + tc1, _ := tcs[1].(map[string]any) + fn1, _ := tc1["function"].(map[string]any) + if args, _ := fn1["arguments"].(string); args != "{}" { + t.Fatalf("bad1 args = %q, want {}", args) + } + tc2, _ := tcs[2].(map[string]any) + fn2, _ := tc2["function"].(map[string]any) + if args, _ := fn2["arguments"].(string); args != "{}" { + t.Fatalf("bad2 args = %q, want {}", args) + } +} + +// TestNormalizeMessages_KeepsValidToolCalls verifies that valid tool_calls +// (correct name and valid JSON object arguments) are preserved unchanged. +func TestNormalizeMessages_KeepsValidToolCalls(t *testing.T) { + validNames := map[string]bool{"bash": true, "read": true} + messages := []map[string]any{ + {"role": "user", "content": "list files"}, + {"role": "assistant", "content": "sure.", "tool_calls": []any{ + map[string]any{"id": "call_1", "type": "function", "function": map[string]any{"name": "bash", "arguments": `{"command":"ls"}`}}, + map[string]any{"id": "call_2", "type": "function", "function": map[string]any{"name": "read", "arguments": `{"file":"a.txt"}`}}, + }}, + {"role": "tool", "tool_call_id": "call_1", "content": "file1"}, + {"role": "tool", "tool_call_id": "call_2", "content": "content"}, + } + result := normalizeMessages(messages, validNames) + if len(result) != 4 { + t.Fatalf("messages = %d, want 4", len(result)) + } + asst := result[1] + tcs, _ := asst["tool_calls"].([]any) + if len(tcs) != 2 { + t.Fatalf("tool_calls = %d items, want 2", len(tcs)) + } + // Verify the first tool_call is unchanged + tc, _ := tcs[0].(map[string]any) + fn, _ := tc["function"].(map[string]any) + if name, _ := fn["name"].(string); name != "bash" { + t.Fatalf("name = %q, want %q", name, "bash") + } + if args, _ := fn["arguments"].(string); args != `{"command":"ls"}` { + t.Fatalf("args = %q, want %q", args, `{"command":"ls"}`) + } +} + +// TestNormalizeMessages_NoToolNamesSkipsNameCheck verifies that when no +// tools are declared (empty map), tool_calls are not name-sanitized +// (only argument validity is checked). +func TestNormalizeMessages_NoToolNamesSkipsNameCheck(t *testing.T) { + messages := []map[string]any{ + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": "ok.", "tool_calls": []any{ + map[string]any{"id": "call_1", "type": "function", "function": map[string]any{"name": "custom_fn", "arguments": `{"x":1}`}}, + }}, + {"role": "tool", "tool_call_id": "call_1", "content": "result"}, + } + result := normalizeMessages(messages, nil) + if len(result) != 3 { + t.Fatalf("messages = %d, want 3", len(result)) + } + asst := result[1] + tcs, _ := asst["tool_calls"].([]any) + tc, _ := tcs[0].(map[string]any) + fn, _ := tc["function"].(map[string]any) + // Name should be unchanged (no tools to validate against) + if name, _ := fn["name"].(string); name != "custom_fn" { + t.Fatalf("name = %q, want %q (should not be changed)", name, "custom_fn") + } +} + +// TestNormalizeMessages_ErrorMdScenario mirrors the exact error.md +// request: an assistant message with 5 tool_calls where 4 have +// name:"invalid" (with args containing {"tool":"task","error":"…"}) and 1 +// has name:"task" with args:"-1". Verifies that after sanitization, all +// 5 calls have valid names and valid JSON object arguments. +func TestNormalizeMessages_ErrorMdScenario(t *testing.T) { + validNames := map[string]bool{"task": true, "bash": true} + messages := []map[string]any{ + {"role": "user", "content": "create tasks"}, + {"role": "assistant", "content": "confirming.", "tool_calls": []any{ + map[string]any{"id": "call_1", "type": "function", "function": map[string]any{"name": "invalid", "arguments": `{"tool":"task","error":"JSON parsing failed"}`}}, + map[string]any{"id": "call_2", "type": "function", "function": map[string]any{"name": "invalid", "arguments": `{"tool":"task","error":"JSON parsing failed"}`}}, + map[string]any{"id": "call_3", "type": "function", "function": map[string]any{"name": "task", "arguments": "-1"}}, + map[string]any{"id": "call_4", "type": "function", "function": map[string]any{"name": "invalid", "arguments": `{"tool":"task","error":"JSON parsing failed"}`}}, + map[string]any{"id": "call_5", "type": "function", "function": map[string]any{"name": "invalid", "arguments": `{"tool":"task","error":"JSON parsing failed"}`}}, + }}, + {"role": "tool", "tool_call_id": "call_1", "content": "error 1"}, + {"role": "tool", "tool_call_id": "call_2", "content": "error 2"}, + {"role": "tool", "tool_call_id": "call_3", "content": "error 3"}, + {"role": "tool", "tool_call_id": "call_4", "content": "error 4"}, + {"role": "tool", "tool_call_id": "call_5", "content": "error 5"}, + } + result := normalizeMessages(messages, validNames) + // All 7 messages preserved (no removal) + if len(result) != 7 { + t.Fatalf("messages = %d, want 7", len(result)) + } + asst := result[1] + tcs, _ := asst["tool_calls"].([]any) + if len(tcs) != 5 { + t.Fatalf("tool_calls = %d items, want 5", len(tcs)) + } + for i, raw := range tcs { + tc, _ := raw.(map[string]any) + fn, _ := tc["function"].(map[string]any) + name, _ := fn["name"].(string) + args, _ := fn["arguments"].(string) + // All names should be "task" (extracted from args or already valid) + if name != "task" { + t.Errorf("tool_call[%d] name = %q, want %q", i, name, "task") + } + // All args should be valid JSON objects + if !isValidJSONObject(args) { + t.Errorf("tool_call[%d] args = %q, not a valid JSON object", i, args) + } + } +} + +// TestIsValidJSONObject verifies the JSON object validation used to filter +// tool_call arguments. +func TestIsValidJSONObject(t *testing.T) { + cases := []struct { + input string + want bool + }{ + {`{"key":"value"}`, true}, + {`{}`, true}, + {`{"nested":{"a":1}}`, true}, + {`-1`, false}, + {`true`, false}, + {`"string"`, false}, + {`[1,2,3]`, false}, + {``, false}, + {` `, false}, + {`{invalid json}`, false}, + } + for _, tc := range cases { + if got := isValidJSONObject(tc.input); got != tc.want { + t.Errorf("isValidJSONObject(%q) = %v, want %v", tc.input, got, tc.want) + } + } +}