From 7a7eab640325a0d454c78b57b30562161ce80fe3 Mon Sep 17 00:00:00 2001 From: ZanzyTHEbar Date: Sat, 21 Feb 2026 22:27:02 +0000 Subject: [PATCH] test: use cmp.Diff for assertions and refactor migrate_sessions tests Replace assert.Equal/require.Equal with assert.Empty(cmp.Diff(...)) across 73 test files. cmp.Diff produces clearer failure output by showing the exact difference between expected and actual values. Refactor pkg/memory/migrate_sessions_test.go: consolidate TestMigrateFileSessions_Basic and related cases into a single table-driven TestMigrateFileSessions with structured test cases (setup, want, extra assertions). Packages affected: eval/go_evals, internal/fantasy/*, pkg/agent, pkg/itr/*, pkg/memory/*, pkg/rlm/*, pkg/runtime, pkg/security/*, pkg/session, pkg/skills, pkg/sync, pkg/tools/*, pkg/worker. --- eval/go_evals/eval_test.go | 17 +- internal/fantasy/agent_stream_test.go | 63 +-- internal/fantasy/agent_test.go | 190 +++++---- .../providers/anthropic/anthropic_test.go | 14 +- .../fantasy/providers/azure/azure_test.go | 3 +- .../fantasy/providers/openai/openai_test.go | 391 +++++++++--------- .../openaicompat/openaicompat_test.go | 28 +- .../fantasy/providertests/anthropic_test.go | 4 +- .../providertests/azure_responses_test.go | 10 +- internal/fantasy/providertests/common_test.go | 4 +- .../providertests/openai_responses_test.go | 4 +- .../fantasy/providertests/openrouter_test.go | 4 +- .../providertests/provider_registry_test.go | 44 +- internal/fantasy/providertests/vercel_test.go | 4 +- internal/fantasy/react_fsm_test.go | 25 +- internal/fantasy/schema/schema_test.go | 76 ++-- internal/fantasy/tool_runtime_dag_test.go | 11 +- .../fantasy/tool_runtime_parallel_test.go | 19 +- internal/fantasy/tool_test.go | 7 +- pkg/agent/kv_delegate_test.go | 13 +- pkg/agent/offloading_runtime_test.go | 7 +- pkg/agent/state_store_test.go | 6 +- pkg/agent/tool_result_search_test.go | 9 +- pkg/itr/commands_test.go | 32 +- pkg/itr/dag/executor_test.go | 13 +- pkg/itr/dag/planner_test.go | 37 +- pkg/itr/dag/replan_test.go | 3 +- pkg/itr/dag/resolver_test.go | 19 +- pkg/itr/dag/router_test.go | 23 +- pkg/itr/wasm/runtime_test.go | 17 +- pkg/itr/wasm/transport_test.go | 3 +- pkg/memory/dag/backfill_test.go | 9 +- pkg/memory/dag/dag_test.go | 45 +- pkg/memory/delegate/sqlite_audit_test.go | 19 +- pkg/memory/delegate/sqlite_dag_test.go | 8 +- pkg/memory/delegate/sqlite_doc_test.go | 17 +- .../delegate/sqlite_integration_test.go | 67 ++- pkg/memory/delegate/sqlite_kv_test.go | 9 +- .../delegate/sqlite_session_msg_test.go | 58 ++- pkg/memory/integration_test.go | 19 +- pkg/memory/kernel_contract_test.go | 3 +- pkg/memory/migrate_sessions_test.go | 311 +++++++------- pkg/memory/observation/observation_test.go | 32 +- pkg/memory/store/chunker_test.go | 9 +- pkg/memory/store/memory_store_test.go | 37 +- pkg/memory/store/memory_tool_test.go | 5 +- pkg/memory/store/queue_test.go | 13 +- pkg/memory/store/scorer_test.go | 11 +- pkg/rlm/engine_test.go | 19 +- pkg/rlm/fanout_test.go | 15 +- pkg/rlm/rope_test.go | 37 +- pkg/rlm/strategy_test.go | 21 +- pkg/runtime/runtime_test.go | 12 +- pkg/security/jsonextract_test.go | 25 +- pkg/security/redact_test.go | 9 +- pkg/security/securebus/audit_test.go | 19 +- pkg/security/securebus/bus_test.go | 15 +- .../securebus/socket_transport_test.go | 5 +- pkg/security/urlguard_test.go | 3 +- pkg/security/vault_test.go | 7 +- pkg/security/zkp_test.go | 21 +- pkg/session/manager_test.go | 12 +- pkg/skills/graph_test.go | 21 +- pkg/skills/templates_test.go | 3 +- pkg/sync/identity_test.go | 15 +- pkg/tools/agentic_map_test.go | 9 +- pkg/tools/focus_test.go | 19 +- pkg/tools/llm_map_test.go | 5 +- pkg/tools/map_runtime_integration_test.go | 19 +- pkg/tools/obligation_test.go | 17 +- pkg/tools/retrieval_test.go | 7 +- pkg/tools/skills_test.go | 7 +- pkg/worker/worker_test.go | 15 +- 73 files changed, 1135 insertions(+), 994 deletions(-) diff --git a/eval/go_evals/eval_test.go b/eval/go_evals/eval_test.go index ca2da32fb..def1650dc 100644 --- a/eval/go_evals/eval_test.go +++ b/eval/go_evals/eval_test.go @@ -9,6 +9,7 @@ import ( "github.com/ZanzyTHEbar/dragonscale/pkg/config" "github.com/ZanzyTHEbar/dragonscale/pkg/tools" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -484,11 +485,11 @@ func TestConfig_DefaultValues(t *testing.T) { cfg := config.DefaultConfig() assert.True(t, cfg.Agents.Defaults.RestrictToSandbox, "restrict to sandbox should be on by default") - assert.Equal(t, 20, cfg.Agents.Defaults.MaxToolIterations, "max tool iterations default") - assert.Equal(t, 0.7, cfg.Agents.Defaults.Temperature, "temperature default") - assert.Equal(t, 8192, cfg.Agents.Defaults.MaxTokens, "max tokens default") - assert.Equal(t, 768, cfg.Memory.EmbeddingDims, "embedding dims default") - assert.Equal(t, 4000, cfg.Memory.OffloadThresholdTokens, "offload threshold default") + assert.Empty(t, cmp.Diff(20, cfg.Agents.Defaults.MaxToolIterations), "max tool iterations default") + assert.Empty(t, cmp.Diff(0.7, cfg.Agents.Defaults.Temperature), "temperature default") + assert.Empty(t, cmp.Diff(8192, cfg.Agents.Defaults.MaxTokens), "max tokens default") + assert.Empty(t, cmp.Diff(768, cfg.Memory.EmbeddingDims), "embedding dims default") + assert.Empty(t, cmp.Diff(4000, cfg.Memory.OffloadThresholdTokens), "offload threshold default") } func TestConfig_LoadEvalConfigs(t *testing.T) { @@ -510,7 +511,7 @@ func TestConfig_LoadEvalConfigs(t *testing.T) { require.NoError(t, err, "config should load without error") assert.True(t, cfg.Agents.Defaults.RestrictToSandbox, "restrict_to_sandbox should always be true for eval") - assert.Equal(t, tc.expectIterations, cfg.Agents.Defaults.MaxToolIterations, "max_tool_iterations") + assert.Empty(t, cmp.Diff(tc.expectIterations, cfg.Agents.Defaults.MaxToolIterations), "max_tool_iterations") }) } } @@ -519,7 +520,7 @@ func TestConfig_MissingFileReturnsDefaults(t *testing.T) { t.Parallel() cfg, err := config.LoadConfig("/nonexistent/path/config.json") require.NoError(t, err, "missing config should return defaults, not error") - assert.Equal(t, 768, cfg.Memory.EmbeddingDims, "should have default embedding dims") + assert.Empty(t, cmp.Diff(768, cfg.Memory.EmbeddingDims), "should have default embedding dims") } // --------------------------------------------------------------------------- @@ -573,7 +574,7 @@ func TestToolSchema_JSONRoundtrip(t *testing.T) { require.NoError(t, err, "schema JSON should parse back") fn := parsed["function"].(map[string]interface{}) - assert.Equal(t, tool.Name(), fn["name"]) + assert.Empty(t, cmp.Diff(tool.Name(), fn["name"])) }) } } diff --git a/internal/fantasy/agent_stream_test.go b/internal/fantasy/agent_stream_test.go index 5f2354db7..810ffdbc6 100644 --- a/internal/fantasy/agent_stream_test.go +++ b/internal/fantasy/agent_stream_test.go @@ -6,7 +6,9 @@ import ( "testing" jsonv2 "github.com/go-json-experiment/json" + "github.com/google/go-cmp/cmp" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -316,30 +318,30 @@ func TestStreamingAgentWithTools(t *testing.T) { Prompt: "Echo 'test'", OnToolInputStart: func(id, toolName string) error { toolInputStartCalled = true - require.Equal(t, "tool-1", id) - require.Equal(t, "echo", toolName) + assert.Empty(t, cmp.Diff("tool-1", id)) + assert.Empty(t, cmp.Diff("echo", toolName)) return nil }, OnToolInputDelta: func(id, delta string) error { toolInputDeltaCalled = true - require.Equal(t, "tool-1", id) + assert.Empty(t, cmp.Diff("tool-1", id)) require.Contains(t, []string{`{"message"`, `: "test"}`}, delta) return nil }, OnToolInputEnd: func(id string) error { toolInputEndCalled = true - require.Equal(t, "tool-1", id) + assert.Empty(t, cmp.Diff("tool-1", id)) return nil }, OnToolCall: func(toolCall ToolCallContent) error { toolCallCalled = true - require.Equal(t, "echo", toolCall.ToolName) - require.Equal(t, `{"message": "test"}`, toolCall.Input) + assert.Empty(t, cmp.Diff("echo", toolCall.ToolName)) + assert.Empty(t, cmp.Diff(`{"message": "test"}`, toolCall.Input)) return nil }, OnToolResult: func(result ToolResultContent) error { toolResultCalled = true - require.Equal(t, "echo", result.ToolName) + assert.Empty(t, cmp.Diff("echo", result.ToolName)) return nil }, } @@ -354,17 +356,17 @@ func TestStreamingAgentWithTools(t *testing.T) { require.True(t, toolInputEndCalled, "OnToolInputEnd should have been called") require.True(t, toolCallCalled, "OnToolCall should have been called") require.True(t, toolResultCalled, "OnToolResult should have been called") - require.Equal(t, 2, len(result.Steps)) // Two steps: tool call + final response + assert.Empty(t, cmp.Diff(2, len(result.Steps))) // Two steps: tool call + final response // Check that tool was executed in first step firstStep := result.Steps[0] toolCalls := firstStep.Content.ToolCalls() - require.Equal(t, 1, len(toolCalls)) - require.Equal(t, "echo", toolCalls[0].ToolName) + assert.Empty(t, cmp.Diff(1, len(toolCalls))) + assert.Empty(t, cmp.Diff("echo", toolCalls[0].ToolName)) toolResults := firstStep.Content.ToolResults() - require.Equal(t, 1, len(toolResults)) - require.Equal(t, "echo", toolResults[0].ToolName) + assert.Empty(t, cmp.Diff(1, len(toolResults))) + assert.Empty(t, cmp.Diff("echo", toolResults[0].ToolName)) } // TestStreamingAgentTextDeltas tests text streaming (mirrors TS textStream tests) @@ -417,11 +419,12 @@ func TestStreamingAgentTextDeltas(t *testing.T) { result, err := agent.Stream(ctx, streamCall) require.NoError(t, err) + assert. - // Verify text deltas match expected pattern - require.Equal(t, []string{"Hello", ", ", "world!"}, textDeltas) - require.Equal(t, "Hello, world!", result.Response.Content.Text()) - require.Equal(t, int64(13), result.TotalUsage.TotalTokens) + // Verify text deltas match expected pattern + Empty(t, cmp.Diff([]string{"Hello", ", ", "world!"}, textDeltas)) + assert.Empty(t, cmp.Diff("Hello, world!", result.Response.Content.Text())) + assert.Empty(t, cmp.Diff(int64(13), result.TotalUsage.TotalTokens)) } // TestStreamingAgentReasoning tests reasoning content (mirrors TS reasoning tests) @@ -481,12 +484,13 @@ func TestStreamingAgentReasoning(t *testing.T) { result, err := agent.Stream(ctx, streamCall) require.NoError(t, err) + assert. - // Verify reasoning and text are separate - require.Equal(t, []string{"I will open the conversation", " with witty banter."}, reasoningDeltas) - require.Equal(t, []string{"Hi there!"}, textDeltas) - require.Equal(t, "Hi there!", result.Response.Content.Text()) - require.Equal(t, "I will open the conversation with witty banter.", result.Response.Content.ReasoningText()) + // Verify reasoning and text are separate + Empty(t, cmp.Diff([]string{"I will open the conversation", " with witty banter."}, reasoningDeltas)) + assert.Empty(t, cmp.Diff([]string{"Hi there!"}, textDeltas)) + assert.Empty(t, cmp.Diff("Hi there!", result.Response.Content.Text())) + assert.Empty(t, cmp.Diff("I will open the conversation with witty banter.", result.Response.Content.ReasoningText())) } // TestStreamingAgentError tests error handling (mirrors TS error tests) @@ -583,16 +587,17 @@ func TestStreamingAgentSources(t *testing.T) { result, err := agent.Stream(ctx, streamCall) require.NoError(t, err) + assert. - // Verify sources were captured - require.Equal(t, 2, len(sources)) - require.Equal(t, SourceTypeURL, sources[0].SourceType) - require.Equal(t, "https://example.com", sources[0].URL) - require.Equal(t, "Example", sources[0].Title) - require.Equal(t, SourceTypeDocument, sources[1].SourceType) - require.Equal(t, "Document Example", sources[1].Title) + // Verify sources were captured + Empty(t, cmp.Diff(2, len(sources))) + assert.Empty(t, cmp.Diff(SourceTypeURL, sources[0].SourceType)) + assert.Empty(t, cmp.Diff("https://example.com", sources[0].URL)) + assert.Empty(t, cmp.Diff("Example", sources[0].Title)) + assert.Empty(t, cmp.Diff(SourceTypeDocument, sources[1].SourceType)) + assert.Empty(t, cmp.Diff("Document Example", sources[1].Title)) // Verify sources are in final result resultSources := result.Response.Content.Sources() - require.Equal(t, 2, len(resultSources)) + assert.Empty(t, cmp.Diff(2, len(resultSources))) } diff --git a/internal/fantasy/agent_test.go b/internal/fantasy/agent_test.go index 6a00c543b..946ab4055 100644 --- a/internal/fantasy/agent_test.go +++ b/internal/fantasy/agent_test.go @@ -7,7 +7,9 @@ import ( "testing" jsonv2 "github.com/go-json-experiment/json" + "github.com/google/go-cmp/cmp" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -104,7 +106,7 @@ func TestAgent_Generate_ResultContent_AllTypes(t *testing.T) { "tool1", "Test tool", func(ctx context.Context, input TestInput, _ ToolCall) (ToolResponse, error) { - require.Equal(t, "value", input.Value) + assert.Empty(t, cmp.Diff("value", input.Value)) return ToolResponse{Content: "result1", IsError: false}, nil }, ) @@ -159,33 +161,33 @@ func TestAgent_Generate_ResultContent_AllTypes(t *testing.T) { // Verify each content type in order textContent, ok := AsContentType[TextContent](result.Response.Content[0]) require.True(t, ok) - require.Equal(t, "Hello, world!", textContent.Text) + assert.Empty(t, cmp.Diff("Hello, world!", textContent.Text)) sourceContent, ok := AsContentType[SourceContent](result.Response.Content[1]) require.True(t, ok) - require.Equal(t, "123", sourceContent.ID) + assert.Empty(t, cmp.Diff("123", sourceContent.ID)) fileContent, ok := AsContentType[FileContent](result.Response.Content[2]) require.True(t, ok) - require.Equal(t, []byte{1, 2, 3}, fileContent.Data) + assert.Empty(t, cmp.Diff([]byte{1, 2, 3}, fileContent.Data)) reasoningContent, ok := AsContentType[ReasoningContent](result.Response.Content[3]) require.True(t, ok) - require.Equal(t, "I will open the conversation with witty banter.", reasoningContent.Text) + assert.Empty(t, cmp.Diff("I will open the conversation with witty banter.", reasoningContent.Text)) toolCallContent, ok := AsContentType[ToolCallContent](result.Response.Content[4]) require.True(t, ok) - require.Equal(t, "call-1", toolCallContent.ToolCallID) + assert.Empty(t, cmp.Diff("call-1", toolCallContent.ToolCallID)) moreTextContent, ok := AsContentType[TextContent](result.Response.Content[5]) require.True(t, ok) - require.Equal(t, "More text", moreTextContent.Text) + assert.Empty(t, cmp.Diff("More text", moreTextContent.Text)) // Tool result should be appended toolResultContent, ok := AsContentType[ToolResultContent](result.Response.Content[6]) require.True(t, ok) - require.Equal(t, "call-1", toolResultContent.ToolCallID) - require.Equal(t, "tool1", toolResultContent.ToolName) + assert.Empty(t, cmp.Diff("call-1", toolResultContent.ToolCallID)) + assert.Empty(t, cmp.Diff("tool1", toolResultContent.ToolName)) } // Test result.text extraction @@ -218,7 +220,7 @@ func TestAgent_Generate_ResultText(t *testing.T) { // Test text extraction from content text := result.Response.Content.Text() - require.Equal(t, "Hello, world!", text) + assert.Empty(t, cmp.Diff("Hello, world!", text)) } // Test result.toolCalls extraction (matches TS test exactly) @@ -254,11 +256,11 @@ func TestAgent_Generate_ResultToolCalls(t *testing.T) { generateFunc: func(ctx context.Context, call Call) (*Response, error) { // Verify tools are passed correctly require.Len(t, call.Tools, 2) - require.Equal(t, ToolChoiceAuto, *call.ToolChoice) // Should be auto, not required + assert.Empty(t, cmp.Diff(ToolChoiceAuto, *call.ToolChoice)) // Should be auto, not required // Verify prompt structure require.Len(t, call.Prompt, 1) - require.Equal(t, MessageRoleUser, call.Prompt[0].Role) + assert.Empty(t, cmp.Diff(MessageRoleUser, call.Prompt[0].Role)) return &Response{ Content: []Content{ @@ -296,14 +298,14 @@ func TestAgent_Generate_ResultToolCalls(t *testing.T) { } require.Len(t, toolCalls, 1) - require.Equal(t, "call-1", toolCalls[0].ToolCallID) - require.Equal(t, "tool1", toolCalls[0].ToolName) + assert.Empty(t, cmp.Diff("call-1", toolCalls[0].ToolCallID)) + assert.Empty(t, cmp.Diff("tool1", toolCalls[0].ToolName)) // Parse and verify input var input map[string]any err = jsonv2.Unmarshal([]byte(toolCalls[0].Input), &input) require.NoError(t, err) - require.Equal(t, "value", input["value"]) + assert.Empty(t, cmp.Diff("value", input["value"])) } // Test result.toolResults extraction (matches TS test exactly) @@ -319,7 +321,7 @@ func TestAgent_Generate_ResultToolResults(t *testing.T) { "tool1", "Test tool", func(ctx context.Context, input TestInput, _ ToolCall) (ToolResponse, error) { - require.Equal(t, "value", input.Value) + assert.Empty(t, cmp.Diff("value", input.Value)) return ToolResponse{Content: "result1", IsError: false}, nil }, ) @@ -328,11 +330,11 @@ func TestAgent_Generate_ResultToolResults(t *testing.T) { generateFunc: func(ctx context.Context, call Call) (*Response, error) { // Verify tools and tool choice require.Len(t, call.Tools, 1) - require.Equal(t, ToolChoiceAuto, *call.ToolChoice) + assert.Empty(t, cmp.Diff(ToolChoiceAuto, *call.ToolChoice)) // Verify prompt require.Len(t, call.Prompt, 1) - require.Equal(t, MessageRoleUser, call.Prompt[0].Role) + assert.Empty(t, cmp.Diff(MessageRoleUser, call.Prompt[0].Role)) return &Response{ Content: []Content{ @@ -370,13 +372,13 @@ func TestAgent_Generate_ResultToolResults(t *testing.T) { } require.Len(t, toolResults, 1) - require.Equal(t, "call-1", toolResults[0].ToolCallID) - require.Equal(t, "tool1", toolResults[0].ToolName) + assert.Empty(t, cmp.Diff("call-1", toolResults[0].ToolCallID)) + assert.Empty(t, cmp.Diff("tool1", toolResults[0].ToolName)) // Verify result content textResult, ok := toolResults[0].Result.(ToolResultOutputContentText) require.True(t, ok) - require.Equal(t, "result1", textResult.Text) + assert.Empty(t, cmp.Diff("result1", textResult.Text)) } // Test multi-step scenario (matches TS "2 steps: initial, tool-result" test) @@ -392,7 +394,7 @@ func TestAgent_Generate_MultipleSteps(t *testing.T) { "tool1", "Test tool", func(ctx context.Context, input TestInput, _ ToolCall) (ToolResponse, error) { - require.Equal(t, "value", input.Value) + assert.Empty(t, cmp.Diff("value", input.Value)) return ToolResponse{Content: "result1", IsError: false}, nil }, ) @@ -447,17 +449,20 @@ func TestAgent_Generate_MultipleSteps(t *testing.T) { require.NoError(t, err) require.NotNil(t, result) require.Len(t, result.Steps, 2) + assert. - // Check total usage sums both steps - require.Equal(t, int64(13), result.TotalUsage.InputTokens) // 10 + 3 - require.Equal(t, int64(15), result.TotalUsage.OutputTokens) // 5 + 10 - require.Equal(t, int64(28), result.TotalUsage.TotalTokens) // 15 + 13 + // Check total usage sums both steps + Empty(t, cmp.Diff(int64(13), result.TotalUsage.InputTokens)) + assert. // 10 + 3 + Empty(t, cmp.Diff(int64(15), result.TotalUsage.OutputTokens)) + assert. // 5 + 10 + Empty(t, cmp.Diff(int64(28), result.TotalUsage.TotalTokens)) // 15 + 13 // Final response should be from last step require.Len(t, result.Response.Content, 1) textContent, ok := AsContentType[TextContent](result.Response.Content[0]) require.True(t, ok) - require.Equal(t, "Hello, world!", textContent.Text) + assert.Empty(t, cmp.Diff("Hello, world!", textContent.Text)) // result.toolCalls should be empty (from last step) var toolCalls []ToolCallContent @@ -511,17 +516,19 @@ func TestAgent_Generate_BasicText(t *testing.T) { require.Len(t, result.Response.Content, 1) textContent, ok := AsContentType[TextContent](result.Response.Content[0]) require.True(t, ok) - require.Equal(t, "Hello, world!", textContent.Text) + assert.Empty(t, cmp.Diff("Hello, world!", textContent.Text)) + assert. - // Check usage - require.Equal(t, int64(3), result.Response.Usage.InputTokens) - require.Equal(t, int64(10), result.Response.Usage.OutputTokens) - require.Equal(t, int64(13), result.Response.Usage.TotalTokens) + // Check usage + Empty(t, cmp.Diff(int64(3), result.Response.Usage.InputTokens)) + assert.Empty(t, cmp.Diff(int64(10), result.Response.Usage.OutputTokens)) + assert.Empty(t, cmp.Diff(int64(13), result.Response.Usage.TotalTokens)) + assert. - // Check total usage - require.Equal(t, int64(3), result.TotalUsage.InputTokens) - require.Equal(t, int64(10), result.TotalUsage.OutputTokens) - require.Equal(t, int64(13), result.TotalUsage.TotalTokens) + // Check total usage + Empty(t, cmp.Diff(int64(3), result.TotalUsage.InputTokens)) + assert.Empty(t, cmp.Diff(int64(10), result.TotalUsage.OutputTokens)) + assert.Empty(t, cmp.Diff(int64(13), result.TotalUsage.TotalTokens)) } // Test empty prompt error @@ -547,13 +554,14 @@ func TestAgent_Generate_WithSystemPrompt(t *testing.T) { model := &mockLanguageModel{ generateFunc: func(ctx context.Context, call Call) (*Response, error) { // Verify system message is included - require.Len(t, call.Prompt, 2) // system + user - require.Equal(t, MessageRoleSystem, call.Prompt[0].Role) - require.Equal(t, MessageRoleUser, call.Prompt[1].Role) + require.Len(t, call.Prompt, 2) + assert. // system + user + Empty(t, cmp.Diff(MessageRoleSystem, call.Prompt[0].Role)) + assert.Empty(t, cmp.Diff(MessageRoleUser, call.Prompt[1].Role)) systemPart, ok := call.Prompt[0].Content[0].(TextPart) require.True(t, ok) - require.Equal(t, "You are a helpful assistant", systemPart.Text) + assert.Empty(t, cmp.Diff("You are a helpful assistant", systemPart.Text)) return &Response{ Content: []Content{ @@ -606,7 +614,7 @@ func TestAgent_Generate_OptionsActiveTools(t *testing.T) { require.Len(t, call.Tools, 1) functionTool, ok := call.Tools[0].(FunctionTool) require.True(t, ok) - require.Equal(t, "tool1", functionTool.Name) + assert.Empty(t, cmp.Diff("tool1", functionTool.Name)) return &Response{ Content: []Content{ @@ -644,46 +652,48 @@ func TestResponseContent_Getters(t *testing.T) { ToolCallContent{ToolCallID: "call1", ToolName: "test_tool", Input: `{"arg": "value"}`}, ToolResultContent{ToolCallID: "call1", ToolName: "test_tool", Result: ToolResultOutputContentText{Text: "result"}}, } + assert. - // Test Text() - require.Equal(t, "Hello world", content.Text()) + // Test Text() + Empty(t, cmp.Diff("Hello world", content.Text())) // Test Reasoning() reasoning := content.Reasoning() require.Len(t, reasoning, 1) - require.Equal(t, "Let me think...", reasoning[0].Text) + assert.Empty(t, cmp.Diff("Let me think...", reasoning[0].Text)) + assert. - // Test ReasoningText() - require.Equal(t, "Let me think...", content.ReasoningText()) + // Test ReasoningText() + Empty(t, cmp.Diff("Let me think...", content.ReasoningText())) // Test Files() files := content.Files() require.Len(t, files, 1) - require.Equal(t, "text/plain", files[0].MediaType) - require.Equal(t, []byte("file data"), files[0].Data) + assert.Empty(t, cmp.Diff("text/plain", files[0].MediaType)) + assert.Empty(t, cmp.Diff([]byte("file data"), files[0].Data)) // Test Sources() sources := content.Sources() require.Len(t, sources, 1) - require.Equal(t, SourceTypeURL, sources[0].SourceType) - require.Equal(t, "https://example.com", sources[0].URL) - require.Equal(t, "Example", sources[0].Title) + assert.Empty(t, cmp.Diff(SourceTypeURL, sources[0].SourceType)) + assert.Empty(t, cmp.Diff("https://example.com", sources[0].URL)) + assert.Empty(t, cmp.Diff("Example", sources[0].Title)) // Test ToolCalls() toolCalls := content.ToolCalls() require.Len(t, toolCalls, 1) - require.Equal(t, "call1", toolCalls[0].ToolCallID) - require.Equal(t, "test_tool", toolCalls[0].ToolName) - require.Equal(t, `{"arg": "value"}`, toolCalls[0].Input) + assert.Empty(t, cmp.Diff("call1", toolCalls[0].ToolCallID)) + assert.Empty(t, cmp.Diff("test_tool", toolCalls[0].ToolName)) + assert.Empty(t, cmp.Diff(`{"arg": "value"}`, toolCalls[0].Input)) // Test ToolResults() toolResults := content.ToolResults() require.Len(t, toolResults, 1) - require.Equal(t, "call1", toolResults[0].ToolCallID) - require.Equal(t, "test_tool", toolResults[0].ToolName) + assert.Empty(t, cmp.Diff("call1", toolResults[0].ToolCallID)) + assert.Empty(t, cmp.Diff("test_tool", toolResults[0].ToolName)) result, ok := AsToolResultOutputType[ToolResultOutputContentText](toolResults[0].Result) require.True(t, ok) - require.Equal(t, "result", result.Text) + assert.Empty(t, cmp.Diff("result", result.Text)) } func TestResponseContent_Getters_Empty(t *testing.T) { @@ -691,9 +701,8 @@ func TestResponseContent_Getters_Empty(t *testing.T) { // Test with empty content content := ResponseContent{} - - require.Equal(t, "", content.Text()) - require.Equal(t, "", content.ReasoningText()) + assert.Empty(t, cmp.Diff("", content.Text())) + assert.Empty(t, cmp.Diff("", content.ReasoningText())) require.Empty(t, content.Reasoning()) require.Empty(t, content.Files()) require.Empty(t, content.Sources()) @@ -715,17 +724,18 @@ func TestResponseContent_Getters_MultipleItems(t *testing.T) { // Test multiple reasoning reasoning := content.Reasoning() require.Len(t, reasoning, 2) - require.Equal(t, "First thought", reasoning[0].Text) - require.Equal(t, "Second thought", reasoning[1].Text) + assert.Empty(t, cmp.Diff("First thought", reasoning[0].Text)) + assert.Empty(t, cmp.Diff("Second thought", reasoning[1].Text)) + assert. - // Test concatenated reasoning text - require.Equal(t, "First thoughtSecond thought", content.ReasoningText()) + // Test concatenated reasoning text + Empty(t, cmp.Diff("First thoughtSecond thought", content.ReasoningText())) // Test multiple files files := content.Files() require.Len(t, files, 2) - require.Equal(t, "text/plain", files[0].MediaType) - require.Equal(t, "image/png", files[1].MediaType) + assert.Empty(t, cmp.Diff("text/plain", files[0].MediaType)) + assert.Empty(t, cmp.Diff("image/png", files[1].MediaType)) } func TestStopConditions(t *testing.T) { @@ -910,8 +920,9 @@ func TestStopConditions_Integration(t *testing.T) { require.NoError(t, err) require.NotNil(t, result) - // Should stop on first condition met (finish reason stop) - require.Equal(t, FinishReasonStop, result.Response.FinishReason) + assert. + // Should stop on first condition met (finish reason stop) + Empty(t, cmp.Diff(FinishReasonStop, result.Response.FinishReason)) }) } @@ -959,7 +970,7 @@ func TestPrepareStep(t *testing.T) { require.NoError(t, err) require.NotNil(t, result) - require.Equal(t, "Modified system prompt for step 0", capturedSystemPrompt) + assert.Empty(t, cmp.Diff("Modified system prompt for step 0", capturedSystemPrompt)) }) t.Run("Tool choice modification", func(t *testing.T) { @@ -997,7 +1008,7 @@ func TestPrepareStep(t *testing.T) { require.NoError(t, err) require.NotNil(t, result) require.NotNil(t, capturedToolChoice) - require.Equal(t, ToolChoiceNone, *capturedToolChoice) + assert.Empty(t, cmp.Diff(ToolChoiceNone, *capturedToolChoice)) }) t.Run("Active tools modification", func(t *testing.T) { @@ -1042,7 +1053,7 @@ func TestPrepareStep(t *testing.T) { require.NoError(t, err) require.NotNil(t, result) require.Len(t, capturedToolNames, 1) - require.Equal(t, "tool2", capturedToolNames[0]) + assert.Empty(t, cmp.Diff("tool2", capturedToolNames[0])) }) t.Run("No tools when DisableAllTools is true", func(t *testing.T) { @@ -1080,7 +1091,7 @@ func TestPrepareStep(t *testing.T) { require.NoError(t, err) require.NotNil(t, result) - require.Equal(t, 0, capturedToolCount) // No tools should be passed + assert.Empty(t, cmp.Diff(0, capturedToolCount)) // No tools should be passed }) t.Run("All fields modified together", func(t *testing.T) { @@ -1140,11 +1151,11 @@ func TestPrepareStep(t *testing.T) { require.NoError(t, err) require.NotNil(t, result) - require.Equal(t, "Step-specific system", capturedSystemPrompt) + assert.Empty(t, cmp.Diff("Step-specific system", capturedSystemPrompt)) require.NotNil(t, capturedToolChoice) - require.Equal(t, SpecificToolChoice("tool1"), *capturedToolChoice) + assert.Empty(t, cmp.Diff(SpecificToolChoice("tool1"), *capturedToolChoice)) require.Len(t, capturedToolNames, 1) - require.Equal(t, "tool1", capturedToolNames[0]) + assert.Empty(t, cmp.Diff("tool1", capturedToolNames[0])) }) t.Run("Nil fields use parent values", func(t *testing.T) { @@ -1201,11 +1212,11 @@ func TestPrepareStep(t *testing.T) { require.NoError(t, err) require.NotNil(t, result) - require.Equal(t, "Parent system", capturedSystemPrompt) + assert.Empty(t, cmp.Diff("Parent system", capturedSystemPrompt)) require.NotNil(t, capturedToolChoice) - require.Equal(t, ToolChoiceAuto, *capturedToolChoice) // Default + assert.Empty(t, cmp.Diff(ToolChoiceAuto, *capturedToolChoice)) // Default require.Len(t, capturedToolNames, 1) - require.Equal(t, "tool1", capturedToolNames[0]) + assert.Empty(t, cmp.Diff("tool1", capturedToolNames[0])) }) t.Run("Empty ActiveTools means all tools", func(t *testing.T) { @@ -1399,8 +1410,9 @@ func TestToolCallRepair(t *testing.T) { // Check that tool call was repaired and is now valid toolCalls := result.Steps[0].Content.ToolCalls() require.Len(t, toolCalls, 1) - require.False(t, toolCalls[0].Invalid) // Should be valid after repair - require.Equal(t, `{"value": "repaired"}`, toolCalls[0].Input) // Should have repaired input + require.False(t, toolCalls[0].Invalid) + assert.Empty( // Should be valid after repair + t, cmp.Diff(`{"value": "repaired"}`, toolCalls[0].Input)) // Should have repaired input }) t.Run("Invalid tool call with failed repair", func(t *testing.T) { @@ -1596,8 +1608,8 @@ func TestAgent_MediaToolResponses(t *testing.T) { mediaResult, ok := toolResults[0].Result.(ToolResultOutputContentMedia) require.True(t, ok, "Expected media result") - require.Equal(t, string(imageData), mediaResult.Data) - require.Equal(t, "image/png", mediaResult.MediaType) + assert.Empty(t, cmp.Diff(string(imageData), mediaResult.Data)) + assert.Empty(t, cmp.Diff("image/png", mediaResult.MediaType)) }) t.Run("Media tool response (audio)", func(t *testing.T) { @@ -1648,8 +1660,8 @@ func TestAgent_MediaToolResponses(t *testing.T) { mediaResult, ok := toolResults[0].Result.(ToolResultOutputContentMedia) require.True(t, ok, "Expected media result") - require.Equal(t, string(audioData), mediaResult.Data) - require.Equal(t, "audio/wav", mediaResult.MediaType) + assert.Empty(t, cmp.Diff(string(audioData), mediaResult.Data)) + assert.Empty(t, cmp.Diff("audio/wav", mediaResult.MediaType)) }) t.Run("Media response with text", func(t *testing.T) { @@ -1702,9 +1714,9 @@ func TestAgent_MediaToolResponses(t *testing.T) { mediaResult, ok := toolResults[0].Result.(ToolResultOutputContentMedia) require.True(t, ok, "Expected media result") - require.Equal(t, string(imageData), mediaResult.Data) - require.Equal(t, "image/png", mediaResult.MediaType) - require.Equal(t, "Screenshot captured successfully", mediaResult.Text) + assert.Empty(t, cmp.Diff(string(imageData), mediaResult.Data)) + assert.Empty(t, cmp.Diff("image/png", mediaResult.MediaType)) + assert.Empty(t, cmp.Diff("Screenshot captured successfully", mediaResult.Text)) }) t.Run("Media response preserves metadata", func(t *testing.T) { @@ -1765,7 +1777,7 @@ func TestAgent_MediaToolResponses(t *testing.T) { var metadata ImageMetadata err = jsonv2.Unmarshal([]byte(toolResults[0].ClientMetadata), &metadata) require.NoError(t, err) - require.Equal(t, 800, metadata.Width) - require.Equal(t, 600, metadata.Height) + assert.Empty(t, cmp.Diff(800, metadata.Width)) + assert.Empty(t, cmp.Diff(600, metadata.Height)) }) } diff --git a/internal/fantasy/providers/anthropic/anthropic_test.go b/internal/fantasy/providers/anthropic/anthropic_test.go index b826882c5..1a70ffcb1 100644 --- a/internal/fantasy/providers/anthropic/anthropic_test.go +++ b/internal/fantasy/providers/anthropic/anthropic_test.go @@ -5,6 +5,8 @@ import ( "testing" "charm.land/fantasy" + "github.com/google/go-cmp/cmp" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -41,7 +43,7 @@ func TestToPrompt_DropsEmptyMessages(t *testing.T) { require.Empty(t, systemBlocks) require.Len(t, messages, 1, "should only have user message, assistant message should be dropped") require.Len(t, warnings, 1) - require.Equal(t, fantasy.CallWarningTypeOther, warnings[0].Type) + assert.Empty(t, cmp.Diff(fantasy.CallWarningTypeOther, warnings[0].Type)) require.Contains(t, warnings[0].Message, "dropping empty assistant message") require.Contains(t, warnings[0].Message, "neither user-facing content nor tool calls") }) @@ -76,9 +78,9 @@ func TestToPrompt_DropsEmptyMessages(t *testing.T) { require.Empty(t, systemBlocks) require.Len(t, messages, 1, "should only have user message, assistant message should be dropped") require.Len(t, warnings, 2) - require.Equal(t, fantasy.CallWarningTypeOther, warnings[0].Type) + assert.Empty(t, cmp.Diff(fantasy.CallWarningTypeOther, warnings[0].Type)) require.Contains(t, warnings[0].Message, "sending reasoning content is disabled") - require.Equal(t, fantasy.CallWarningTypeOther, warnings[1].Type) + assert.Empty(t, cmp.Diff(fantasy.CallWarningTypeOther, warnings[1].Type)) require.Contains(t, warnings[1].Message, "dropping empty assistant message") }) @@ -103,7 +105,7 @@ func TestToPrompt_DropsEmptyMessages(t *testing.T) { require.Empty(t, systemBlocks) require.Len(t, messages, 1, "should only have user message") require.Len(t, warnings, 1) - require.Equal(t, fantasy.CallWarningTypeOther, warnings[0].Type) + assert.Empty(t, cmp.Diff(fantasy.CallWarningTypeOther, warnings[0].Type)) require.Contains(t, warnings[0].Message, "dropping empty assistant message") }) @@ -188,7 +190,7 @@ func TestToPrompt_DropsEmptyMessages(t *testing.T) { require.Empty(t, systemBlocks) require.Len(t, messages, 1, "should only have user message") require.Len(t, warnings, 1) - require.Equal(t, fantasy.CallWarningTypeOther, warnings[0].Type) + assert.Empty(t, cmp.Diff(fantasy.CallWarningTypeOther, warnings[0].Type)) require.Contains(t, warnings[0].Message, "dropping empty assistant message") }) @@ -267,7 +269,7 @@ func TestToPrompt_DropsEmptyMessages(t *testing.T) { require.Empty(t, systemBlocks) require.Empty(t, messages) require.Len(t, warnings, 1) - require.Equal(t, fantasy.CallWarningTypeOther, warnings[0].Type) + assert.Empty(t, cmp.Diff(fantasy.CallWarningTypeOther, warnings[0].Type)) require.Contains(t, warnings[0].Message, "dropping empty user message") require.Contains(t, warnings[0].Message, "neither user-facing content nor tool results") }) diff --git a/internal/fantasy/providers/azure/azure_test.go b/internal/fantasy/providers/azure/azure_test.go index d02aa51b3..7edef40f7 100644 --- a/internal/fantasy/providers/azure/azure_test.go +++ b/internal/fantasy/providers/azure/azure_test.go @@ -3,6 +3,7 @@ package azure import ( "testing" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" ) @@ -93,7 +94,7 @@ func TestParseAzureURL(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { result := parseAzureURL(tt.input) - assert.Equal(t, tt.expected, result) + assert.Empty(t, cmp.Diff(tt.expected, result)) }) } } diff --git a/internal/fantasy/providers/openai/openai_test.go b/internal/fantasy/providers/openai/openai_test.go index 121313f80..4a5a140b3 100644 --- a/internal/fantasy/providers/openai/openai_test.go +++ b/internal/fantasy/providers/openai/openai_test.go @@ -10,7 +10,9 @@ import ( "charm.land/fantasy" jsonv2 "github.com/go-json-experiment/json" + "github.com/google/go-cmp/cmp" "github.com/openai/openai-go/v2/packages/param" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -36,7 +38,7 @@ func TestToOpenAiPrompt_SystemMessages(t *testing.T) { systemMsg := messages[0].OfSystem require.NotNil(t, systemMsg) - require.Equal(t, "You are a helpful assistant.", systemMsg.Content.OfString.Value) + assert.Empty(t, cmp.Diff("You are a helpful assistant.", systemMsg.Content.OfString.Value)) }) t.Run("should handle empty system messages", func(t *testing.T) { @@ -76,7 +78,7 @@ func TestToOpenAiPrompt_SystemMessages(t *testing.T) { systemMsg := messages[0].OfSystem require.NotNil(t, systemMsg) - require.Equal(t, "You are a helpful assistant.\nBe concise.", systemMsg.Content.OfString.Value) + assert.Empty(t, cmp.Diff("You are a helpful assistant.\nBe concise.", systemMsg.Content.OfString.Value)) }) } @@ -102,7 +104,7 @@ func TestToOpenAiPrompt_UserMessages(t *testing.T) { userMsg := messages[0].OfUser require.NotNil(t, userMsg) - require.Equal(t, "Hello", userMsg.Content.OfString.Value) + assert.Empty(t, cmp.Diff("Hello", userMsg.Content.OfString.Value)) }) t.Run("should convert messages with image parts", func(t *testing.T) { @@ -136,13 +138,13 @@ func TestToOpenAiPrompt_UserMessages(t *testing.T) { // Check text part textPart := content[0].OfText require.NotNil(t, textPart) - require.Equal(t, "Hello", textPart.Text) + assert.Empty(t, cmp.Diff("Hello", textPart.Text)) // Check image part imagePart := content[1].OfImageURL require.NotNil(t, imagePart) expectedURL := "data:image/png;base64," + base64.StdEncoding.EncodeToString(imageData) - require.Equal(t, expectedURL, imagePart.ImageURL.URL) + assert.Empty(t, cmp.Diff(expectedURL, imagePart.ImageURL.URL)) }) t.Run("should add image detail when specified through provider options", func(t *testing.T) { @@ -177,7 +179,7 @@ func TestToOpenAiPrompt_UserMessages(t *testing.T) { imagePart := content[0].OfImageURL require.NotNil(t, imagePart) - require.Equal(t, "low", imagePart.ImageURL.Detail) + assert.Empty(t, cmp.Diff("low", imagePart.ImageURL.Detail)) }) } @@ -236,8 +238,8 @@ func TestToOpenAiPrompt_FileParts(t *testing.T) { audioPart := content[0].OfInputAudio require.NotNil(t, audioPart) - require.Equal(t, base64.StdEncoding.EncodeToString(audioData), audioPart.InputAudio.Data) - require.Equal(t, "wav", audioPart.InputAudio.Format) + assert.Empty(t, cmp.Diff(base64.StdEncoding.EncodeToString(audioData), audioPart.InputAudio.Data)) + assert.Empty(t, cmp.Diff("wav", audioPart.InputAudio.Format)) }) t.Run("should add audio content for audio/mpeg file parts", func(t *testing.T) { @@ -265,7 +267,7 @@ func TestToOpenAiPrompt_FileParts(t *testing.T) { content := userMsg.Content.OfArrayOfContentParts audioPart := content[0].OfInputAudio require.NotNil(t, audioPart) - require.Equal(t, "mp3", audioPart.InputAudio.Format) + assert.Empty(t, cmp.Diff("mp3", audioPart.InputAudio.Format)) }) t.Run("should add audio content for audio/mp3 file parts", func(t *testing.T) { @@ -293,7 +295,7 @@ func TestToOpenAiPrompt_FileParts(t *testing.T) { content := userMsg.Content.OfArrayOfContentParts audioPart := content[0].OfInputAudio require.NotNil(t, audioPart) - require.Equal(t, "mp3", audioPart.InputAudio.Format) + assert.Empty(t, cmp.Diff("mp3", audioPart.InputAudio.Format)) }) t.Run("should convert messages with PDF file parts", func(t *testing.T) { @@ -324,10 +326,10 @@ func TestToOpenAiPrompt_FileParts(t *testing.T) { filePart := content[0].OfFile require.NotNil(t, filePart) - require.Equal(t, "document.pdf", filePart.File.Filename.Value) + assert.Empty(t, cmp.Diff("document.pdf", filePart.File.Filename.Value)) expectedData := "data:application/pdf;base64," + base64.StdEncoding.EncodeToString(pdfData) - require.Equal(t, expectedData, filePart.File.FileData.Value) + assert.Empty(t, cmp.Diff(expectedData, filePart.File.FileData.Value)) }) t.Run("should convert messages with binary PDF file parts", func(t *testing.T) { @@ -358,7 +360,7 @@ func TestToOpenAiPrompt_FileParts(t *testing.T) { require.NotNil(t, filePart) expectedData := "data:application/pdf;base64," + base64.StdEncoding.EncodeToString(pdfData) - require.Equal(t, expectedData, filePart.File.FileData.Value) + assert.Empty(t, cmp.Diff(expectedData, filePart.File.FileData.Value)) }) t.Run("should convert messages with PDF file parts using file_id", func(t *testing.T) { @@ -385,7 +387,7 @@ func TestToOpenAiPrompt_FileParts(t *testing.T) { content := userMsg.Content.OfArrayOfContentParts filePart := content[0].OfFile require.NotNil(t, filePart) - require.Equal(t, "file-pdf-12345", filePart.File.FileID.Value) + assert.Empty(t, cmp.Diff("file-pdf-12345", filePart.File.FileID.Value)) require.True(t, param.IsOmitted(filePart.File.FileData)) require.True(t, param.IsOmitted(filePart.File.Filename)) }) @@ -415,7 +417,7 @@ func TestToOpenAiPrompt_FileParts(t *testing.T) { content := userMsg.Content.OfArrayOfContentParts filePart := content[0].OfFile require.NotNil(t, filePart) - require.Equal(t, "part-0.pdf", filePart.File.Filename.Value) + assert.Empty(t, cmp.Diff("part-0.pdf", filePart.File.Filename.Value)) }) } @@ -463,20 +465,20 @@ func TestToOpenAiPrompt_ToolCalls(t *testing.T) { // Check assistant message with tool call assistantMsg := messages[0].OfAssistant require.NotNil(t, assistantMsg) - require.Equal(t, "", assistantMsg.Content.OfString.Value) + assert.Empty(t, cmp.Diff("", assistantMsg.Content.OfString.Value)) require.Len(t, assistantMsg.ToolCalls, 1) toolCall := assistantMsg.ToolCalls[0].OfFunction require.NotNil(t, toolCall) - require.Equal(t, "quux", toolCall.ID) - require.Equal(t, "thwomp", toolCall.Function.Name) - require.Equal(t, string(inputJSON), toolCall.Function.Arguments) + assert.Empty(t, cmp.Diff("quux", toolCall.ID)) + assert.Empty(t, cmp.Diff("thwomp", toolCall.Function.Name)) + assert.Empty(t, cmp.Diff(string(inputJSON), toolCall.Function.Arguments)) // Check tool message toolMsg := messages[1].OfTool require.NotNil(t, toolMsg) - require.Equal(t, string(outputJSON), toolMsg.Content.OfString.Value) - require.Equal(t, "quux", toolMsg.ToolCallID) + assert.Empty(t, cmp.Diff(string(outputJSON), toolMsg.Content.OfString.Value)) + assert.Empty(t, cmp.Diff("quux", toolMsg.ToolCallID)) }) t.Run("should handle different tool output types", func(t *testing.T) { @@ -510,14 +512,14 @@ func TestToOpenAiPrompt_ToolCalls(t *testing.T) { // Check first tool message (text) textToolMsg := messages[0].OfTool require.NotNil(t, textToolMsg) - require.Equal(t, "Hello world", textToolMsg.Content.OfString.Value) - require.Equal(t, "text-tool", textToolMsg.ToolCallID) + assert.Empty(t, cmp.Diff("Hello world", textToolMsg.Content.OfString.Value)) + assert.Empty(t, cmp.Diff("text-tool", textToolMsg.ToolCallID)) // Check second tool message (error) errorToolMsg := messages[1].OfTool require.NotNil(t, errorToolMsg) - require.Equal(t, "Something went wrong", errorToolMsg.Content.OfString.Value) - require.Equal(t, "error-tool", errorToolMsg.ToolCallID) + assert.Empty(t, cmp.Diff("Something went wrong", errorToolMsg.Content.OfString.Value)) + assert.Empty(t, cmp.Diff("error-tool", errorToolMsg.ToolCallID)) }) } @@ -543,7 +545,7 @@ func TestToOpenAiPrompt_AssistantMessages(t *testing.T) { assistantMsg := messages[0].OfAssistant require.NotNil(t, assistantMsg) - require.Equal(t, "Hello, how can I help you?", assistantMsg.Content.OfString.Value) + assert.Empty(t, cmp.Diff("Hello, how can I help you?", assistantMsg.Content.OfString.Value)) }) t.Run("should handle assistant messages with mixed content", func(t *testing.T) { @@ -573,13 +575,13 @@ func TestToOpenAiPrompt_AssistantMessages(t *testing.T) { assistantMsg := messages[0].OfAssistant require.NotNil(t, assistantMsg) - require.Equal(t, "Let me search for that.", assistantMsg.Content.OfString.Value) + assert.Empty(t, cmp.Diff("Let me search for that.", assistantMsg.Content.OfString.Value)) require.Len(t, assistantMsg.ToolCalls, 1) toolCall := assistantMsg.ToolCalls[0].OfFunction - require.Equal(t, "call-123", toolCall.ID) - require.Equal(t, "search", toolCall.Function.Name) - require.Equal(t, string(inputJSON), toolCall.Function.Arguments) + assert.Empty(t, cmp.Diff("call-123", toolCall.ID)) + assert.Empty(t, cmp.Diff("search", toolCall.Function.Name)) + assert.Empty(t, cmp.Diff(string(inputJSON), toolCall.Function.Arguments)) }) } @@ -825,7 +827,7 @@ func TestDoGenerate(t *testing.T) { textContent, ok := result.Content[0].(fantasy.TextContent) require.True(t, ok) - require.Equal(t, "Hello, World!", textContent.Text) + assert.Empty(t, cmp.Diff("Hello, World!", textContent.Text)) }) t.Run("should extract usage", func(t *testing.T) { @@ -854,9 +856,9 @@ func TestDoGenerate(t *testing.T) { }) require.NoError(t, err) - require.Equal(t, int64(20), result.Usage.InputTokens) - require.Equal(t, int64(5), result.Usage.OutputTokens) - require.Equal(t, int64(25), result.Usage.TotalTokens) + assert.Empty(t, cmp.Diff(int64(20), result.Usage.InputTokens)) + assert.Empty(t, cmp.Diff(int64(5), result.Usage.OutputTokens)) + assert.Empty(t, cmp.Diff(int64(25), result.Usage.TotalTokens)) }) t.Run("should send request body", func(t *testing.T) { @@ -882,17 +884,17 @@ func TestDoGenerate(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "POST", call.method) - require.Equal(t, "/chat/completions", call.path) - require.Equal(t, "gpt-3.5-turbo", call.body["model"]) + assert.Empty(t, cmp.Diff("POST", call.method)) + assert.Empty(t, cmp.Diff("/chat/completions", call.path)) + assert.Empty(t, cmp.Diff("gpt-3.5-turbo", call.body["model"])) messages, ok := call.body["messages"].([]any) require.True(t, ok) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) }) t.Run("should support partial usage", func(t *testing.T) { @@ -920,9 +922,9 @@ func TestDoGenerate(t *testing.T) { }) require.NoError(t, err) - require.Equal(t, int64(20), result.Usage.InputTokens) - require.Equal(t, int64(0), result.Usage.OutputTokens) - require.Equal(t, int64(20), result.Usage.TotalTokens) + assert.Empty(t, cmp.Diff(int64(20), result.Usage.InputTokens)) + assert.Empty(t, cmp.Diff(int64(0), result.Usage.OutputTokens)) + assert.Empty(t, cmp.Diff(int64(20), result.Usage.TotalTokens)) }) t.Run("should extract logprobs", func(t *testing.T) { @@ -982,7 +984,7 @@ func TestDoGenerate(t *testing.T) { }) require.NoError(t, err) - require.Equal(t, fantasy.FinishReasonStop, result.FinishReason) + assert.Empty(t, cmp.Diff(fantasy.FinishReasonStop, result.FinishReason)) }) t.Run("should support unknown finish reason", func(t *testing.T) { @@ -1007,7 +1009,7 @@ func TestDoGenerate(t *testing.T) { }) require.NoError(t, err) - require.Equal(t, fantasy.FinishReasonUnknown, result.FinishReason) + assert.Empty(t, cmp.Diff(fantasy.FinishReasonUnknown, result.FinishReason)) }) t.Run("should pass the model and the messages", func(t *testing.T) { @@ -1035,14 +1037,14 @@ func TestDoGenerate(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "gpt-3.5-turbo", call.body["model"]) + assert.Empty(t, cmp.Diff("gpt-3.5-turbo", call.body["model"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) }) t.Run("should pass settings", func(t *testing.T) { @@ -1075,15 +1077,15 @@ func TestDoGenerate(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "gpt-3.5-turbo", call.body["model"]) + assert.Empty(t, cmp.Diff("gpt-3.5-turbo", call.body["model"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) logitBias := call.body["logit_bias"].(map[string]any) - require.Equal(t, float64(-100), logitBias["50256"]) - require.Equal(t, false, call.body["parallel_tool_calls"]) - require.Equal(t, "test-user-id", call.body["user"]) + assert.Empty(t, cmp.Diff(float64(-100), logitBias["50256"])) + assert.Empty(t, cmp.Diff(false, call.body["parallel_tool_calls"])) + assert.Empty(t, cmp.Diff("test-user-id", call.body["user"])) }) t.Run("should pass reasoningEffort setting", func(t *testing.T) { @@ -1116,15 +1118,15 @@ func TestDoGenerate(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "o1-mini", call.body["model"]) - require.Equal(t, "low", call.body["reasoning_effort"]) + assert.Empty(t, cmp.Diff("o1-mini", call.body["model"])) + assert.Empty(t, cmp.Diff("low", call.body["reasoning_effort"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) }) t.Run("should pass textVerbosity setting", func(t *testing.T) { @@ -1155,15 +1157,15 @@ func TestDoGenerate(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "gpt-4o", call.body["model"]) - require.Equal(t, "low", call.body["verbosity"]) + assert.Empty(t, cmp.Diff("gpt-4o", call.body["model"])) + assert.Empty(t, cmp.Diff("low", call.body["verbosity"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) }) t.Run("should pass tools and toolChoice", func(t *testing.T) { @@ -1208,7 +1210,7 @@ func TestDoGenerate(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "gpt-3.5-turbo", call.body["model"]) + assert.Empty(t, cmp.Diff("gpt-3.5-turbo", call.body["model"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) @@ -1217,17 +1219,17 @@ func TestDoGenerate(t *testing.T) { require.Len(t, tools, 1) tool := tools[0].(map[string]any) - require.Equal(t, "function", tool["type"]) + assert.Empty(t, cmp.Diff("function", tool["type"])) function := tool["function"].(map[string]any) - require.Equal(t, "test-tool", function["name"]) - require.Equal(t, false, function["strict"]) + assert.Empty(t, cmp.Diff("test-tool", function["name"])) + assert.Empty(t, cmp.Diff(false, function["strict"])) toolChoice := call.body["tool_choice"].(map[string]any) - require.Equal(t, "function", toolChoice["type"]) + assert.Empty(t, cmp.Diff("function", toolChoice["type"])) toolChoiceFunction := toolChoice["function"].(map[string]any) - require.Equal(t, "test-tool", toolChoiceFunction["name"]) + assert.Empty(t, cmp.Diff("test-tool", toolChoiceFunction["name"])) }) t.Run("should parse tool results", func(t *testing.T) { @@ -1282,9 +1284,9 @@ func TestDoGenerate(t *testing.T) { toolCall, ok := result.Content[0].(fantasy.ToolCallContent) require.True(t, ok) - require.Equal(t, "call_O17Uplv4lJvD6DVdIvFFeRMw", toolCall.ToolCallID) - require.Equal(t, "test-tool", toolCall.ToolName) - require.Equal(t, `{"value":"Spark"}`, toolCall.Input) + assert.Empty(t, cmp.Diff("call_O17Uplv4lJvD6DVdIvFFeRMw", toolCall.ToolCallID)) + assert.Empty(t, cmp.Diff("test-tool", toolCall.ToolName)) + assert.Empty(t, cmp.Diff(`{"value":"Spark"}`, toolCall.Input)) }) t.Run("should handle ToolChoiceRequired", func(t *testing.T) { @@ -1329,21 +1331,21 @@ func TestDoGenerate(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "gpt-3.5-turbo", call.body["model"]) + assert.Empty(t, cmp.Diff("gpt-3.5-turbo", call.body["model"])) // Verify tool is present tools := call.body["tools"].([]any) require.Len(t, tools, 1) tool := tools[0].(map[string]any) - require.Equal(t, "function", tool["type"]) + assert.Empty(t, cmp.Diff("function", tool["type"])) function := tool["function"].(map[string]any) - require.Equal(t, "test-tool", function["name"]) + assert.Empty(t, cmp.Diff("test-tool", function["name"])) // Verify tool_choice is set to "required" (not a function name) toolChoice := call.body["tool_choice"] - require.Equal(t, "required", toolChoice) + assert.Empty(t, cmp.Diff("required", toolChoice)) }) t.Run("should parse annotations/citations", func(t *testing.T) { @@ -1383,13 +1385,13 @@ func TestDoGenerate(t *testing.T) { textContent, ok := result.Content[0].(fantasy.TextContent) require.True(t, ok) - require.Equal(t, "Based on the search results [doc1], I found information.", textContent.Text) + assert.Empty(t, cmp.Diff("Based on the search results [doc1], I found information.", textContent.Text)) sourceContent, ok := result.Content[1].(fantasy.SourceContent) require.True(t, ok) - require.Equal(t, fantasy.SourceTypeURL, sourceContent.SourceType) - require.Equal(t, "https://example.com/doc1.pdf", sourceContent.URL) - require.Equal(t, "Document 1", sourceContent.Title) + assert.Empty(t, cmp.Diff(fantasy.SourceTypeURL, sourceContent.SourceType)) + assert.Empty(t, cmp.Diff("https://example.com/doc1.pdf", sourceContent.URL)) + assert.Empty(t, cmp.Diff("Document 1", sourceContent.Title)) require.NotEmpty(t, sourceContent.ID) }) @@ -1422,10 +1424,10 @@ func TestDoGenerate(t *testing.T) { }) require.NoError(t, err) - require.Equal(t, int64(1152), result.Usage.CacheReadTokens) - require.Equal(t, int64(15), result.Usage.InputTokens) - require.Equal(t, int64(20), result.Usage.OutputTokens) - require.Equal(t, int64(35), result.Usage.TotalTokens) + assert.Empty(t, cmp.Diff(int64(1152), result.Usage.CacheReadTokens)) + assert.Empty(t, cmp.Diff(int64(15), result.Usage.InputTokens)) + assert.Empty(t, cmp.Diff(int64(20), result.Usage.OutputTokens)) + assert.Empty(t, cmp.Diff(int64(35), result.Usage.TotalTokens)) }) t.Run("should return accepted_prediction_tokens and rejected_prediction_tokens", func(t *testing.T) { @@ -1463,8 +1465,8 @@ func TestDoGenerate(t *testing.T) { openaiMeta, ok := result.ProviderMetadata["openai"].(*ProviderMetadata) require.True(t, ok) - require.Equal(t, int64(123), openaiMeta.AcceptedPredictionTokens) - require.Equal(t, int64(456), openaiMeta.RejectedPredictionTokens) + assert.Empty(t, cmp.Diff(int64(123), openaiMeta.AcceptedPredictionTokens)) + assert.Empty(t, cmp.Diff(int64(456), openaiMeta.RejectedPredictionTokens)) }) t.Run("should clear out temperature, top_p, frequency_penalty, presence_penalty for reasoning models", func(t *testing.T) { @@ -1494,14 +1496,14 @@ func TestDoGenerate(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "o1-preview", call.body["model"]) + assert.Empty(t, cmp.Diff("o1-preview", call.body["model"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) // These should not be present require.Nil(t, call.body["temperature"]) @@ -1511,8 +1513,8 @@ func TestDoGenerate(t *testing.T) { // Should have warnings require.Len(t, result.Warnings, 4) - require.Equal(t, fantasy.CallWarningTypeUnsupportedSetting, result.Warnings[0].Type) - require.Equal(t, "temperature", result.Warnings[0].Setting) + assert.Empty(t, cmp.Diff(fantasy.CallWarningTypeUnsupportedSetting, result.Warnings[0].Type)) + assert.Empty(t, cmp.Diff("temperature", result.Warnings[0].Setting)) require.Contains(t, result.Warnings[0].Details, "temperature is not supported for reasoning models") }) @@ -1540,16 +1542,16 @@ func TestDoGenerate(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "o1-preview", call.body["model"]) - require.Equal(t, float64(1000), call.body["max_completion_tokens"]) + assert.Empty(t, cmp.Diff("o1-preview", call.body["model"])) + assert.Empty(t, cmp.Diff(float64(1000), call.body["max_completion_tokens"])) require.Nil(t, call.body["max_tokens"]) messages := call.body["messages"].([]any) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) }) t.Run("should return reasoning tokens", func(t *testing.T) { @@ -1581,10 +1583,10 @@ func TestDoGenerate(t *testing.T) { }) require.NoError(t, err) - require.Equal(t, int64(15), result.Usage.InputTokens) - require.Equal(t, int64(20), result.Usage.OutputTokens) - require.Equal(t, int64(35), result.Usage.TotalTokens) - require.Equal(t, int64(10), result.Usage.ReasoningTokens) + assert.Empty(t, cmp.Diff(int64(15), result.Usage.InputTokens)) + assert.Empty(t, cmp.Diff(int64(20), result.Usage.OutputTokens)) + assert.Empty(t, cmp.Diff(int64(35), result.Usage.TotalTokens)) + assert.Empty(t, cmp.Diff(int64(10), result.Usage.ReasoningTokens)) }) t.Run("should send max_completion_tokens extension setting", func(t *testing.T) { @@ -1615,15 +1617,15 @@ func TestDoGenerate(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "o1-preview", call.body["model"]) - require.Equal(t, float64(255), call.body["max_completion_tokens"]) + assert.Empty(t, cmp.Diff("o1-preview", call.body["model"])) + assert.Empty(t, cmp.Diff(float64(255), call.body["max_completion_tokens"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) }) t.Run("should send prediction extension setting", func(t *testing.T) { @@ -1657,18 +1659,18 @@ func TestDoGenerate(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "gpt-3.5-turbo", call.body["model"]) + assert.Empty(t, cmp.Diff("gpt-3.5-turbo", call.body["model"])) prediction := call.body["prediction"].(map[string]any) - require.Equal(t, "content", prediction["type"]) - require.Equal(t, "Hello, World!", prediction["content"]) + assert.Empty(t, cmp.Diff("content", prediction["type"])) + assert.Empty(t, cmp.Diff("Hello, World!", prediction["content"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) }) t.Run("should send store extension setting", func(t *testing.T) { @@ -1699,15 +1701,15 @@ func TestDoGenerate(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "gpt-3.5-turbo", call.body["model"]) - require.Equal(t, true, call.body["store"]) + assert.Empty(t, cmp.Diff("gpt-3.5-turbo", call.body["model"])) + assert.Empty(t, cmp.Diff(true, call.body["store"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) }) t.Run("should send metadata extension values", func(t *testing.T) { @@ -1740,17 +1742,17 @@ func TestDoGenerate(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "gpt-3.5-turbo", call.body["model"]) + assert.Empty(t, cmp.Diff("gpt-3.5-turbo", call.body["model"])) metadata := call.body["metadata"].(map[string]any) - require.Equal(t, "value", metadata["custom"]) + assert.Empty(t, cmp.Diff("value", metadata["custom"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) }) t.Run("should send promptCacheKey extension value", func(t *testing.T) { @@ -1781,15 +1783,15 @@ func TestDoGenerate(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "gpt-3.5-turbo", call.body["model"]) - require.Equal(t, "test-cache-key-123", call.body["prompt_cache_key"]) + assert.Empty(t, cmp.Diff("gpt-3.5-turbo", call.body["model"])) + assert.Empty(t, cmp.Diff("test-cache-key-123", call.body["prompt_cache_key"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) }) t.Run("should send safety_identifier extension value", func(t *testing.T) { @@ -1820,15 +1822,15 @@ func TestDoGenerate(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "gpt-3.5-turbo", call.body["model"]) - require.Equal(t, "test-safety-identifier-123", call.body["safety_identifier"]) + assert.Empty(t, cmp.Diff("gpt-3.5-turbo", call.body["model"])) + assert.Empty(t, cmp.Diff("test-safety-identifier-123", call.body["safety_identifier"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) }) t.Run("should remove temperature setting for search preview models", func(t *testing.T) { @@ -1855,12 +1857,12 @@ func TestDoGenerate(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "gpt-4o-search-preview", call.body["model"]) + assert.Empty(t, cmp.Diff("gpt-4o-search-preview", call.body["model"])) require.Nil(t, call.body["temperature"]) require.Len(t, result.Warnings, 1) - require.Equal(t, fantasy.CallWarningTypeUnsupportedSetting, result.Warnings[0].Type) - require.Equal(t, "temperature", result.Warnings[0].Setting) + assert.Empty(t, cmp.Diff(fantasy.CallWarningTypeUnsupportedSetting, result.Warnings[0].Type)) + assert.Empty(t, cmp.Diff("temperature", result.Warnings[0].Setting)) require.Contains(t, result.Warnings[0].Details, "search preview models") }) @@ -1892,15 +1894,15 @@ func TestDoGenerate(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "o3-mini", call.body["model"]) - require.Equal(t, "flex", call.body["service_tier"]) + assert.Empty(t, cmp.Diff("o3-mini", call.body["model"])) + assert.Empty(t, cmp.Diff("flex", call.body["service_tier"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) }) t.Run("should show warning when using flex processing with unsupported model", func(t *testing.T) { @@ -1932,8 +1934,8 @@ func TestDoGenerate(t *testing.T) { require.Nil(t, call.body["service_tier"]) require.Len(t, result.Warnings, 1) - require.Equal(t, fantasy.CallWarningTypeUnsupportedSetting, result.Warnings[0].Type) - require.Equal(t, "ServiceTier", result.Warnings[0].Setting) + assert.Empty(t, cmp.Diff(fantasy.CallWarningTypeUnsupportedSetting, result.Warnings[0].Type)) + assert.Empty(t, cmp.Diff("ServiceTier", result.Warnings[0].Setting)) require.Contains(t, result.Warnings[0].Details, "flex processing is only available") }) @@ -1963,15 +1965,15 @@ func TestDoGenerate(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "gpt-4o-mini", call.body["model"]) - require.Equal(t, "priority", call.body["service_tier"]) + assert.Empty(t, cmp.Diff("gpt-4o-mini", call.body["model"])) + assert.Empty(t, cmp.Diff("priority", call.body["service_tier"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) }) t.Run("should show warning when using priority processing with unsupported model", func(t *testing.T) { @@ -2003,8 +2005,8 @@ func TestDoGenerate(t *testing.T) { require.Nil(t, call.body["service_tier"]) require.Len(t, result.Warnings, 1) - require.Equal(t, fantasy.CallWarningTypeUnsupportedSetting, result.Warnings[0].Type) - require.Equal(t, "ServiceTier", result.Warnings[0].Setting) + assert.Empty(t, cmp.Diff(fantasy.CallWarningTypeUnsupportedSetting, result.Warnings[0].Type)) + assert.Empty(t, cmp.Diff("ServiceTier", result.Warnings[0].Setting)) require.Contains(t, result.Warnings[0].Details, "priority processing is only available") }) } @@ -2329,14 +2331,14 @@ func TestDoStream(t *testing.T) { require.NotEqual(t, -1, textStart) require.NotEqual(t, -1, textEnd) require.NotEqual(t, -1, finish) - require.Equal(t, []string{"Hello", ", ", "World!"}, deltas) + assert.Empty(t, cmp.Diff([]string{"Hello", ", ", "World!"}, deltas)) // Check finish part finishPart := parts[finish] - require.Equal(t, fantasy.FinishReasonStop, finishPart.FinishReason) - require.Equal(t, int64(17), finishPart.Usage.InputTokens) - require.Equal(t, int64(227), finishPart.Usage.OutputTokens) - require.Equal(t, int64(244), finishPart.Usage.TotalTokens) + assert.Empty(t, cmp.Diff(fantasy.FinishReasonStop, finishPart.FinishReason)) + assert.Empty(t, cmp.Diff(int64(17), finishPart.Usage.InputTokens)) + assert.Empty(t, cmp.Diff(int64(227), finishPart.Usage.OutputTokens)) + assert.Empty(t, cmp.Diff(int64(244), finishPart.Usage.TotalTokens)) }) t.Run("should stream tool deltas", func(t *testing.T) { @@ -2387,17 +2389,17 @@ func TestDoStream(t *testing.T) { switch part.Type { case fantasy.StreamPartTypeToolInputStart: toolInputStart = i - require.Equal(t, "call_O17Uplv4lJvD6DVdIvFFeRMw", part.ID) - require.Equal(t, "test-tool", part.ToolCallName) + assert.Empty(t, cmp.Diff("call_O17Uplv4lJvD6DVdIvFFeRMw", part.ID)) + assert.Empty(t, cmp.Diff("test-tool", part.ToolCallName)) case fantasy.StreamPartTypeToolInputDelta: toolDeltas = append(toolDeltas, part.Delta) case fantasy.StreamPartTypeToolInputEnd: toolInputEnd = i case fantasy.StreamPartTypeToolCall: toolCall = i - require.Equal(t, "call_O17Uplv4lJvD6DVdIvFFeRMw", part.ID) - require.Equal(t, "test-tool", part.ToolCallName) - require.Equal(t, `{"value":"Sparkle Day"}`, part.ToolCallInput) + assert.Empty(t, cmp.Diff("call_O17Uplv4lJvD6DVdIvFFeRMw", part.ID)) + assert.Empty(t, cmp.Diff("test-tool", part.ToolCallName)) + assert.Empty(t, cmp.Diff(`{"value":"Sparkle Day"}`, part.ToolCallInput)) } } @@ -2410,7 +2412,7 @@ func TestDoStream(t *testing.T) { for _, delta := range toolDeltas { fullInput.WriteString(delta) } - require.Equal(t, `{"value":"Sparkle Day"}`, fullInput.String()) + assert.Empty(t, cmp.Diff(`{"value":"Sparkle Day"}`, fullInput.String())) }) t.Run("should stream annotations/citations", func(t *testing.T) { @@ -2460,9 +2462,9 @@ func TestDoStream(t *testing.T) { } require.NotNil(t, sourcePart) - require.Equal(t, fantasy.SourceTypeURL, sourcePart.SourceType) - require.Equal(t, "https://example.com/doc1.pdf", sourcePart.URL) - require.Equal(t, "Document 1", sourcePart.Title) + assert.Empty(t, cmp.Diff(fantasy.SourceTypeURL, sourcePart.SourceType)) + assert.Empty(t, cmp.Diff("https://example.com/doc1.pdf", sourcePart.URL)) + assert.Empty(t, cmp.Diff("Document 1", sourcePart.Title)) require.NotEmpty(t, sourcePart.ID) }) @@ -2531,20 +2533,20 @@ func TestDoStream(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "POST", call.method) - require.Equal(t, "/chat/completions", call.path) - require.Equal(t, "gpt-3.5-turbo", call.body["model"]) - require.Equal(t, true, call.body["stream"]) + assert.Empty(t, cmp.Diff("POST", call.method)) + assert.Empty(t, cmp.Diff("/chat/completions", call.path)) + assert.Empty(t, cmp.Diff("gpt-3.5-turbo", call.body["model"])) + assert.Empty(t, cmp.Diff(true, call.body["stream"])) streamOptions := call.body["stream_options"].(map[string]any) - require.Equal(t, true, streamOptions["include_usage"]) + assert.Empty(t, cmp.Diff(true, streamOptions["include_usage"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) }) t.Run("should return cached tokens in providerMetadata", func(t *testing.T) { @@ -2591,10 +2593,10 @@ func TestDoStream(t *testing.T) { } require.NotNil(t, finishPart) - require.Equal(t, int64(1152), finishPart.Usage.CacheReadTokens) - require.Equal(t, int64(15), finishPart.Usage.InputTokens) - require.Equal(t, int64(20), finishPart.Usage.OutputTokens) - require.Equal(t, int64(35), finishPart.Usage.TotalTokens) + assert.Empty(t, cmp.Diff(int64(1152), finishPart.Usage.CacheReadTokens)) + assert.Empty(t, cmp.Diff(int64(15), finishPart.Usage.InputTokens)) + assert.Empty(t, cmp.Diff(int64(20), finishPart.Usage.OutputTokens)) + assert.Empty(t, cmp.Diff(int64(35), finishPart.Usage.TotalTokens)) }) t.Run("should return accepted_prediction_tokens and rejected_prediction_tokens", func(t *testing.T) { @@ -2646,8 +2648,8 @@ func TestDoStream(t *testing.T) { openaiMeta, ok := finishPart.ProviderMetadata["openai"].(*ProviderMetadata) require.True(t, ok) - require.Equal(t, int64(123), openaiMeta.AcceptedPredictionTokens) - require.Equal(t, int64(456), openaiMeta.RejectedPredictionTokens) + assert.Empty(t, cmp.Diff(int64(123), openaiMeta.AcceptedPredictionTokens)) + assert.Empty(t, cmp.Diff(int64(456), openaiMeta.RejectedPredictionTokens)) }) t.Run("should send store extension setting", func(t *testing.T) { @@ -2678,19 +2680,19 @@ func TestDoStream(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "gpt-3.5-turbo", call.body["model"]) - require.Equal(t, true, call.body["stream"]) - require.Equal(t, true, call.body["store"]) + assert.Empty(t, cmp.Diff("gpt-3.5-turbo", call.body["model"])) + assert.Empty(t, cmp.Diff(true, call.body["stream"])) + assert.Empty(t, cmp.Diff(true, call.body["store"])) streamOptions := call.body["stream_options"].(map[string]any) - require.Equal(t, true, streamOptions["include_usage"]) + assert.Empty(t, cmp.Diff(true, streamOptions["include_usage"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) }) t.Run("should send metadata extension values", func(t *testing.T) { @@ -2723,21 +2725,21 @@ func TestDoStream(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "gpt-3.5-turbo", call.body["model"]) - require.Equal(t, true, call.body["stream"]) + assert.Empty(t, cmp.Diff("gpt-3.5-turbo", call.body["model"])) + assert.Empty(t, cmp.Diff(true, call.body["stream"])) metadata := call.body["metadata"].(map[string]any) - require.Equal(t, "value", metadata["custom"]) + assert.Empty(t, cmp.Diff("value", metadata["custom"])) streamOptions := call.body["stream_options"].(map[string]any) - require.Equal(t, true, streamOptions["include_usage"]) + assert.Empty(t, cmp.Diff(true, streamOptions["include_usage"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) }) t.Run("should send serviceTier flex processing setting in streaming", func(t *testing.T) { @@ -2768,19 +2770,19 @@ func TestDoStream(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "o3-mini", call.body["model"]) - require.Equal(t, "flex", call.body["service_tier"]) - require.Equal(t, true, call.body["stream"]) + assert.Empty(t, cmp.Diff("o3-mini", call.body["model"])) + assert.Empty(t, cmp.Diff("flex", call.body["service_tier"])) + assert.Empty(t, cmp.Diff(true, call.body["stream"])) streamOptions := call.body["stream_options"].(map[string]any) - require.Equal(t, true, streamOptions["include_usage"]) + assert.Empty(t, cmp.Diff(true, streamOptions["include_usage"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) }) t.Run("should send serviceTier priority processing setting in streaming", func(t *testing.T) { @@ -2811,19 +2813,19 @@ func TestDoStream(t *testing.T) { require.Len(t, server.calls, 1) call := server.calls[0] - require.Equal(t, "gpt-4o-mini", call.body["model"]) - require.Equal(t, "priority", call.body["service_tier"]) - require.Equal(t, true, call.body["stream"]) + assert.Empty(t, cmp.Diff("gpt-4o-mini", call.body["model"])) + assert.Empty(t, cmp.Diff("priority", call.body["service_tier"])) + assert.Empty(t, cmp.Diff(true, call.body["stream"])) streamOptions := call.body["stream_options"].(map[string]any) - require.Equal(t, true, streamOptions["include_usage"]) + assert.Empty(t, cmp.Diff(true, streamOptions["include_usage"])) messages := call.body["messages"].([]any) require.Len(t, messages, 1) message := messages[0].(map[string]any) - require.Equal(t, "user", message["role"]) - require.Equal(t, "Hello", message["content"]) + assert.Empty(t, cmp.Diff("user", message["role"])) + assert.Empty(t, cmp.Diff("Hello", message["content"])) }) t.Run("should stream text delta for reasoning models", func(t *testing.T) { @@ -2860,9 +2862,10 @@ func TestDoStream(t *testing.T) { textDeltas = append(textDeltas, part.Delta) } } + assert. - // Should contain the text content (without empty delta) - require.Equal(t, []string{"Hello, World!"}, textDeltas) + // Should contain the text content (without empty delta) + Empty(t, cmp.Diff([]string{"Hello, World!"}, textDeltas)) }) t.Run("should send reasoning tokens", func(t *testing.T) { @@ -2910,10 +2913,10 @@ func TestDoStream(t *testing.T) { } require.NotNil(t, finishPart) - require.Equal(t, int64(15), finishPart.Usage.InputTokens) - require.Equal(t, int64(20), finishPart.Usage.OutputTokens) - require.Equal(t, int64(35), finishPart.Usage.TotalTokens) - require.Equal(t, int64(10), finishPart.Usage.ReasoningTokens) + assert.Empty(t, cmp.Diff(int64(15), finishPart.Usage.InputTokens)) + assert.Empty(t, cmp.Diff(int64(20), finishPart.Usage.OutputTokens)) + assert.Empty(t, cmp.Diff(int64(35), finishPart.Usage.TotalTokens)) + assert.Empty(t, cmp.Diff(int64(10), finishPart.Usage.ReasoningTokens)) }) } @@ -2940,7 +2943,7 @@ func TestDefaultToPrompt_DropsEmptyMessages(t *testing.T) { require.Len(t, messages, 1, "should only have user message") require.Len(t, warnings, 1) - require.Equal(t, fantasy.CallWarningTypeOther, warnings[0].Type) + assert.Empty(t, cmp.Diff(fantasy.CallWarningTypeOther, warnings[0].Type)) require.Contains(t, warnings[0].Message, "dropping empty assistant message") }) @@ -3105,7 +3108,7 @@ func TestResponsesToPrompt_DropsEmptyMessages(t *testing.T) { require.Len(t, input, 1, "should only have user message") require.Len(t, warnings, 1) - require.Equal(t, fantasy.CallWarningTypeOther, warnings[0].Type) + assert.Empty(t, cmp.Diff(fantasy.CallWarningTypeOther, warnings[0].Type)) require.Contains(t, warnings[0].Message, "dropping empty assistant message") }) diff --git a/internal/fantasy/providers/openaicompat/openaicompat_test.go b/internal/fantasy/providers/openaicompat/openaicompat_test.go index a603e5e66..fc728b4c5 100644 --- a/internal/fantasy/providers/openaicompat/openaicompat_test.go +++ b/internal/fantasy/providers/openaicompat/openaicompat_test.go @@ -5,6 +5,8 @@ import ( "testing" "charm.land/fantasy" + "github.com/google/go-cmp/cmp" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -44,22 +46,22 @@ func TestToPromptFunc_ReasoningContent(t *testing.T) { // First message (user) - no reasoning msg1 := messages[0].OfUser require.NotNil(t, msg1) - require.Equal(t, "What is 2+2?", msg1.Content.OfString.Value) + assert.Empty(t, cmp.Diff("What is 2+2?", msg1.Content.OfString.Value)) // Second message (assistant) - with reasoning msg2 := messages[1].OfAssistant require.NotNil(t, msg2) - require.Equal(t, "The answer is 4.", msg2.Content.OfString.Value) + assert.Empty(t, cmp.Diff("The answer is 4.", msg2.Content.OfString.Value)) // Check reasoning_content in extra fields extraFields := msg2.ExtraFields() reasoningContent, hasReasoning := extraFields["reasoning_content"] require.True(t, hasReasoning) - require.Equal(t, "Let me think... 2+2 equals 4.", reasoningContent) + assert.Empty(t, cmp.Diff("Let me think... 2+2 equals 4.", reasoningContent)) // Third message (user) - no reasoning msg3 := messages[2].OfUser require.NotNil(t, msg3) - require.Equal(t, "What about 3+3?", msg3.Content.OfString.Value) + assert.Empty(t, cmp.Diff("What about 3+3?", msg3.Content.OfString.Value)) }) t.Run("should handle assistant messages with only reasoning content", func(t *testing.T) { @@ -89,7 +91,7 @@ func TestToPromptFunc_ReasoningContent(t *testing.T) { // User message - unchanged msg := messages[0].OfUser require.NotNil(t, msg) - require.Equal(t, "Hello", msg.Content.OfString.Value) + assert.Empty(t, cmp.Diff("Hello", msg.Content.OfString.Value)) }) t.Run("should not add reasoning_content to messages without reasoning", func(t *testing.T) { @@ -118,7 +120,7 @@ func TestToPromptFunc_ReasoningContent(t *testing.T) { // Assistant message without reasoning msg := messages[1].OfAssistant require.NotNil(t, msg) - require.Equal(t, "Hi there!", msg.Content.OfString.Value) + assert.Empty(t, cmp.Diff("Hi there!", msg.Content.OfString.Value)) extraFields := msg.ExtraFields() _, hasReasoning := extraFields["reasoning_content"] require.False(t, hasReasoning) @@ -150,12 +152,12 @@ func TestToPromptFunc_ReasoningContent(t *testing.T) { // System message - unchanged systemMsg := messages[0].OfSystem require.NotNil(t, systemMsg) - require.Equal(t, "You are helpful.", systemMsg.Content.OfString.Value) + assert.Empty(t, cmp.Diff("You are helpful.", systemMsg.Content.OfString.Value)) // User message - unchanged userMsg := messages[1].OfUser require.NotNil(t, userMsg) - require.Equal(t, "Hello", userMsg.Content.OfString.Value) + assert.Empty(t, cmp.Diff("Hello", userMsg.Content.OfString.Value)) }) t.Run("should use last assistant TextPart only", func(t *testing.T) { @@ -186,7 +188,7 @@ func TestToPromptFunc_ReasoningContent(t *testing.T) { // Assistant message should use only the last TextPart (matching openai behavior) assistantMsg := messages[1].OfAssistant require.NotNil(t, assistantMsg) - require.Equal(t, "Third part.", assistantMsg.Content.OfString.Value) + assert.Empty(t, cmp.Diff("Third part.", assistantMsg.Content.OfString.Value)) }) t.Run("should include user messages with only unsupported attachments", func(t *testing.T) { @@ -226,11 +228,11 @@ func TestToPromptFunc_ReasoningContent(t *testing.T) { msg1 := messages[0].OfUser require.NotNil(t, msg1) - require.Equal(t, "Hello", msg1.Content.OfString.Value) + assert.Empty(t, cmp.Diff("Hello", msg1.Content.OfString.Value)) msg2 := messages[1].OfUser require.NotNil(t, msg2) - require.Equal(t, "After unsupported", msg2.Content.OfString.Value) + assert.Empty(t, cmp.Diff("After unsupported", msg2.Content.OfString.Value)) }) t.Run("should detect PDF file IDs using strings.HasPrefix", func(t *testing.T) { @@ -264,7 +266,7 @@ func TestToPromptFunc_ReasoningContent(t *testing.T) { // Second content part should be file with file_id filePart := content[1].OfFile require.NotNil(t, filePart) - require.Equal(t, "file-abc123xyz", filePart.File.FileID.Value) + assert.Empty(t, cmp.Diff("file-abc123xyz", filePart.File.FileID.Value)) }) } @@ -291,7 +293,7 @@ func TestToPromptFunc_DropsEmptyMessages(t *testing.T) { require.Len(t, messages, 1, "should only have user message") require.Len(t, warnings, 1) - require.Equal(t, fantasy.CallWarningTypeOther, warnings[0].Type) + assert.Empty(t, cmp.Diff(fantasy.CallWarningTypeOther, warnings[0].Type)) require.Contains(t, warnings[0].Message, "dropping empty assistant message") }) diff --git a/internal/fantasy/providertests/anthropic_test.go b/internal/fantasy/providertests/anthropic_test.go index a41be58a0..ff6eccba7 100644 --- a/internal/fantasy/providertests/anthropic_test.go +++ b/internal/fantasy/providertests/anthropic_test.go @@ -9,6 +9,8 @@ import ( "charm.land/fantasy" "charm.land/fantasy/providers/anthropic" "charm.land/x/vcr" + "github.com/google/go-cmp/cmp" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -145,7 +147,7 @@ func testAnthropicThinking(t *testing.T, result *fantasy.AgentResult) { } require.Greater(t, reasoningContentCount, 0) require.Greater(t, signaturesCount, 0) - require.Equal(t, reasoningContentCount, signaturesCount) + assert.Empty(t, cmp.Diff(reasoningContentCount, signaturesCount)) } func anthropicBuilder(model string) builderFunc { diff --git a/internal/fantasy/providertests/azure_responses_test.go b/internal/fantasy/providertests/azure_responses_test.go index ab4de1eea..11cefed8d 100644 --- a/internal/fantasy/providertests/azure_responses_test.go +++ b/internal/fantasy/providertests/azure_responses_test.go @@ -1,7 +1,7 @@ package providertests import ( - "cmp" + stdcmp "cmp" "net/http" "os" "testing" @@ -10,6 +10,8 @@ import ( "charm.land/fantasy/providers/azure" "charm.land/fantasy/providers/openai" "charm.land/x/vcr" + "github.com/google/go-cmp/cmp" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -29,8 +31,8 @@ func TestAzureResponsesCommon(t *testing.T) { func azureReasoningBuilder(model string) builderFunc { return func(t *testing.T, r *vcr.Recorder) (fantasy.LanguageModel, error) { provider, err := azure.New( - azure.WithBaseURL(cmp.Or(os.Getenv("FANTASY_AZURE_BASE_URL"), defaultBaseURL)), - azure.WithAPIKey(cmp.Or(os.Getenv("FANTASY_AZURE_API_KEY"), "(missing)")), + azure.WithBaseURL(stdcmp.Or(os.Getenv("FANTASY_AZURE_BASE_URL"), defaultBaseURL)), + azure.WithAPIKey(stdcmp.Or(os.Getenv("FANTASY_AZURE_API_KEY"), "(missing)")), azure.WithHTTPClient(&http.Client{Transport: r}), azure.WithUseResponsesAPI(), ) @@ -96,5 +98,5 @@ func testAzureResponsesThinkingWithSummaryThinking(t *testing.T, result *fantasy } require.Greater(t, reasoningContentCount, 0) require.Greater(t, encryptedData, 0) - require.Equal(t, reasoningContentCount, encryptedData) + assert.Empty(t, cmp.Diff(reasoningContentCount, encryptedData)) } diff --git a/internal/fantasy/providertests/common_test.go b/internal/fantasy/providertests/common_test.go index e5f9c1c0e..3abc77f57 100644 --- a/internal/fantasy/providertests/common_test.go +++ b/internal/fantasy/providertests/common_test.go @@ -9,7 +9,9 @@ import ( "charm.land/fantasy" "charm.land/x/vcr" + "github.com/google/go-cmp/cmp" "github.com/joho/godotenv" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -118,7 +120,7 @@ func testTool(t *testing.T, pair builderPair) { require.False(t, tc.Invalid) } require.Len(t, toolCalls, 1) - require.Equal(t, toolCalls[0].ToolName, "weather") + assert.Empty(t, cmp.Diff(toolCalls[0].ToolName, "weather")) want1 := "Florence" want2 := "40" diff --git a/internal/fantasy/providertests/openai_responses_test.go b/internal/fantasy/providertests/openai_responses_test.go index f774e6a28..723ef2015 100644 --- a/internal/fantasy/providertests/openai_responses_test.go +++ b/internal/fantasy/providertests/openai_responses_test.go @@ -8,6 +8,8 @@ import ( "charm.land/fantasy" "charm.land/fantasy/providers/openai" "charm.land/x/vcr" + "github.com/google/go-cmp/cmp" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -95,5 +97,5 @@ func testOpenAIResponsesThinkingWithSummaryThinking(t *testing.T, result *fantas } require.Greater(t, reasoningContentCount, 0) require.Greater(t, encryptedData, 0) - require.Equal(t, reasoningContentCount, encryptedData) + assert.Empty(t, cmp.Diff(reasoningContentCount, encryptedData)) } diff --git a/internal/fantasy/providertests/openrouter_test.go b/internal/fantasy/providertests/openrouter_test.go index b9466a47d..33a4c503e 100644 --- a/internal/fantasy/providertests/openrouter_test.go +++ b/internal/fantasy/providertests/openrouter_test.go @@ -9,6 +9,8 @@ import ( "charm.land/fantasy/providers/anthropic" "charm.land/fantasy/providers/openrouter" "charm.land/x/vcr" + "github.com/google/go-cmp/cmp" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -99,7 +101,7 @@ func testOpenrouterThinkingWithSignature(t *testing.T, result *fantasy.AgentResu } require.Greater(t, reasoningContentCount, 0) require.Greater(t, signaturesCount, 0) - require.Equal(t, reasoningContentCount, signaturesCount) + assert.Empty(t, cmp.Diff(reasoningContentCount, signaturesCount)) // we also add the anthropic metadata so test that testAnthropicThinking(t, result) } diff --git a/internal/fantasy/providertests/provider_registry_test.go b/internal/fantasy/providertests/provider_registry_test.go index bc9e2c632..a4d98632b 100644 --- a/internal/fantasy/providertests/provider_registry_test.go +++ b/internal/fantasy/providertests/provider_registry_test.go @@ -4,6 +4,7 @@ import ( "testing" jsonv2 "github.com/go-json-experiment/json" + "github.com/google/go-cmp/cmp" "charm.land/fantasy" "charm.land/fantasy/providers/anthropic" @@ -11,6 +12,7 @@ import ( "charm.land/fantasy/providers/openai" "charm.land/fantasy/providers/openaicompat" "charm.land/fantasy/providers/openrouter" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -36,11 +38,11 @@ func TestProviderRegistry_Serialization_OpenAIOptions(t *testing.T) { po, ok := raw.ProviderOptions[openai.Name] require.True(t, ok) - require.Equal(t, openai.TypeProviderOptions, po["type"]) // no magic strings + assert.Empty(t, cmp.Diff(openai.TypeProviderOptions, po["type"])) // no magic strings // ensure inner data has the field we set inner, ok := po["data"].(map[string]any) require.True(t, ok) - require.Equal(t, "tester", inner["user"]) + assert.Empty(t, cmp.Diff("tester", inner["user"])) var decoded fantasy.Message require.NoError(t, jsonv2.Unmarshal(data, &decoded)) @@ -50,7 +52,7 @@ func TestProviderRegistry_Serialization_OpenAIOptions(t *testing.T) { opt, ok := got.(*openai.ProviderOptions) require.True(t, ok) require.NotNil(t, opt.User) - require.Equal(t, "tester", *opt.User) + assert.Empty(t, cmp.Diff("tester", *opt.User)) } func TestProviderRegistry_Serialization_OpenAIResponses(t *testing.T) { @@ -81,11 +83,11 @@ func TestProviderRegistry_Serialization_OpenAIResponses(t *testing.T) { require.NoError(t, jsonv2.Unmarshal(data, &raw)) po := raw.ProviderOptions[openai.Name] - require.Equal(t, openai.TypeResponsesProviderOptions, po["type"]) // no magic strings + assert.Empty(t, cmp.Diff(openai.TypeResponsesProviderOptions, po["type"])) // no magic strings inner, ok := po["data"].(map[string]any) require.True(t, ok) - require.Equal(t, "cache-key-1", inner["prompt_cache_key"]) - require.Equal(t, true, inner["parallel_tool_calls"]) + assert.Empty(t, cmp.Diff("cache-key-1", inner["prompt_cache_key"])) + assert.Empty(t, cmp.Diff(true, inner["parallel_tool_calls"])) // Unmarshal back and assert concrete type var decoded fantasy.Message @@ -94,9 +96,9 @@ func TestProviderRegistry_Serialization_OpenAIResponses(t *testing.T) { reqOpts, ok := got.(*openai.ResponsesProviderOptions) require.True(t, ok) require.NotNil(t, reqOpts.PromptCacheKey) - require.Equal(t, "cache-key-1", *reqOpts.PromptCacheKey) + assert.Empty(t, cmp.Diff("cache-key-1", *reqOpts.PromptCacheKey)) require.NotNil(t, reqOpts.ParallelToolCalls) - require.Equal(t, true, *reqOpts.ParallelToolCalls) + assert.Empty(t, cmp.Diff(true, *reqOpts.ParallelToolCalls)) } func TestProviderRegistry_Serialization_OpenAIResponsesReasoningMetadata(t *testing.T) { @@ -132,10 +134,10 @@ func TestProviderRegistry_Serialization_OpenAIResponsesReasoningMetadata(t *test require.True(t, ok) om, ok := pm[openai.Name].(map[string]any) require.True(t, ok) - require.Equal(t, openai.TypeResponsesReasoningMetadata, om["type"]) // no magic strings + assert.Empty(t, cmp.Diff(openai.TypeResponsesReasoningMetadata, om["type"])) // no magic strings inner, ok := om["data"].(map[string]any) require.True(t, ok) - require.Equal(t, "item-123", inner["item_id"]) + assert.Empty(t, cmp.Diff("item-123", inner["item_id"])) // Unmarshal back var decoded fantasy.Response @@ -145,8 +147,8 @@ func TestProviderRegistry_Serialization_OpenAIResponsesReasoningMetadata(t *test require.True(t, ok) meta, ok := val.(*openai.ResponsesReasoningMetadata) require.True(t, ok) - require.Equal(t, "item-123", meta.ItemID) - require.Equal(t, []string{"part1", "part2"}, meta.Summary) + assert.Empty(t, cmp.Diff("item-123", meta.ItemID)) + assert.Empty(t, cmp.Diff([]string{"part1", "part2"}, meta.Summary)) } func TestProviderRegistry_Serialization_AnthropicOptions(t *testing.T) { @@ -175,7 +177,7 @@ func TestProviderRegistry_Serialization_AnthropicOptions(t *testing.T) { opt, ok := got.(*anthropic.ProviderOptions) require.True(t, ok) require.NotNil(t, opt.SendReasoning) - require.Equal(t, true, *opt.SendReasoning) + assert.Empty(t, cmp.Diff(true, *opt.SendReasoning)) } func TestProviderRegistry_Serialization_GoogleOptions(t *testing.T) { @@ -203,8 +205,8 @@ func TestProviderRegistry_Serialization_GoogleOptions(t *testing.T) { require.True(t, ok) opt, ok := got.(*google.ProviderOptions) require.True(t, ok) - require.Equal(t, "cached-123", opt.CachedContent) - require.Equal(t, "BLOCK_ONLY_HIGH", opt.Threshold) + assert.Empty(t, cmp.Diff("cached-123", opt.CachedContent)) + assert.Empty(t, cmp.Diff("BLOCK_ONLY_HIGH", opt.Threshold)) } func TestProviderRegistry_Serialization_OpenRouterOptions(t *testing.T) { @@ -234,9 +236,9 @@ func TestProviderRegistry_Serialization_OpenRouterOptions(t *testing.T) { opt, ok := got.(*openrouter.ProviderOptions) require.True(t, ok) require.NotNil(t, opt.IncludeUsage) - require.Equal(t, true, *opt.IncludeUsage) + assert.Empty(t, cmp.Diff(true, *opt.IncludeUsage)) require.NotNil(t, opt.User) - require.Equal(t, "test-user", *opt.User) + assert.Empty(t, cmp.Diff("test-user", *opt.User)) } func TestProviderRegistry_Serialization_OpenAICompatOptions(t *testing.T) { @@ -266,9 +268,9 @@ func TestProviderRegistry_Serialization_OpenAICompatOptions(t *testing.T) { opt, ok := got.(*openaicompat.ProviderOptions) require.True(t, ok) require.NotNil(t, opt.User) - require.Equal(t, "test-user", *opt.User) + assert.Empty(t, cmp.Diff("test-user", *opt.User)) require.NotNil(t, opt.ReasoningEffort) - require.Equal(t, openai.ReasoningEffortHigh, *opt.ReasoningEffort) + assert.Empty(t, cmp.Diff(openai.ReasoningEffortHigh, *opt.ReasoningEffort)) } func TestProviderRegistry_MultiProvider(t *testing.T) { @@ -301,14 +303,14 @@ func TestProviderRegistry_MultiProvider(t *testing.T) { require.True(t, ok) openaiData, ok := openaiOpt.(*openai.ProviderOptions) require.True(t, ok) - require.Equal(t, "user1", *openaiData.User) + assert.Empty(t, cmp.Diff("user1", *openaiData.User)) // Check Anthropic options anthropicOpt, ok := decoded.ProviderOptions[anthropic.Name] require.True(t, ok) anthropicData, ok := anthropicOpt.(*anthropic.ProviderOptions) require.True(t, ok) - require.Equal(t, true, *anthropicData.SendReasoning) + assert.Empty(t, cmp.Diff(true, *anthropicData.SendReasoning)) } func TestProviderRegistry_ErrorHandling(t *testing.T) { diff --git a/internal/fantasy/providertests/vercel_test.go b/internal/fantasy/providertests/vercel_test.go index 21e73d660..f3e20c75f 100644 --- a/internal/fantasy/providertests/vercel_test.go +++ b/internal/fantasy/providertests/vercel_test.go @@ -9,6 +9,8 @@ import ( "charm.land/fantasy/providers/anthropic" "charm.land/fantasy/providers/vercel" "charm.land/x/vcr" + "github.com/google/go-cmp/cmp" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -94,7 +96,7 @@ func testVercelThinkingWithSignature(t *testing.T, result *fantasy.AgentResult) } require.Greater(t, reasoningContentCount, 0) require.Greater(t, signaturesCount, 0) - require.Equal(t, reasoningContentCount, signaturesCount) + assert.Empty(t, cmp.Diff(reasoningContentCount, signaturesCount)) // we also add the anthropic metadata so test that testAnthropicThinking(t, result) } diff --git a/internal/fantasy/react_fsm_test.go b/internal/fantasy/react_fsm_test.go index a2918fd47..bea514c59 100644 --- a/internal/fantasy/react_fsm_test.go +++ b/internal/fantasy/react_fsm_test.go @@ -5,6 +5,7 @@ import ( "sync" "testing" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -56,9 +57,9 @@ func TestFSM_Start(t *testing.T) { transitions := obs.Snapshot() require.Len(t, transitions, 1) - assert.Equal(t, ReActStateInit, transitions[0].From) - assert.Equal(t, ReActStatePrepareStep, transitions[0].To) - assert.Equal(t, ReActTriggerStart, transitions[0].Trigger) + assert.Empty(t, cmp.Diff(ReActStateInit, transitions[0].From)) + assert.Empty(t, cmp.Diff(ReActStatePrepareStep, transitions[0].To)) + assert.Empty(t, cmp.Diff(ReActTriggerStart, transitions[0].Trigger)) } // TestFSM_FullHappyPath drives one complete step through all states and @@ -76,8 +77,8 @@ func TestFSM_FullHappyPath(t *testing.T) { transitions := obs.Snapshot() // Expected: Init->PrepareStep, PrepareStep->LLM, LLM->Validate, Validate->Execute, Execute->Append, Append->Stop, Stop->Done require.Len(t, transitions, 7) - assert.Equal(t, ReActStateInit, transitions[0].From) - assert.Equal(t, ReActStateDone, transitions[6].To) + assert.Empty(t, cmp.Diff(ReActStateInit, transitions[0].From)) + assert.Empty(t, cmp.Diff(ReActStateDone, transitions[6].To)) } // TestFSM_Continue verifies that the loop can re-enter PrepareStep after a @@ -96,10 +97,10 @@ func TestFSM_Continue(t *testing.T) { transitions := obs.Snapshot() // Two full loops: each has 5 states + Start + Continue + Finished = 13 - assert.Equal(t, 13, len(transitions)) + assert.Empty(t, cmp.Diff(13, len(transitions))) // Second loop re-enters PrepareStep - assert.Equal(t, ReActStatePrepareStep, transitions[6].To) + assert.Empty(t, cmp.Diff(ReActStatePrepareStep, transitions[6].To)) } // TestFSM_StopConditionMet verifies the alternative Done path. @@ -114,7 +115,7 @@ func TestFSM_StopConditionMet(t *testing.T) { f.Fire(ctx, ReActTriggerStopConditionMet) last := obs.Snapshot() - assert.Equal(t, ReActStateDone, last[len(last)-1].To) + assert.Empty(t, cmp.Diff(ReActStateDone, last[len(last)-1].To)) } // TestFSM_ErrorTransition verifies the error state is reachable from any state. @@ -130,7 +131,7 @@ func TestFSM_ErrorTransition(t *testing.T) { transitions := obs.Snapshot() last := transitions[len(transitions)-1] - assert.Equal(t, ReActStateError, last.To) + assert.Empty(t, cmp.Diff(ReActStateError, last.To)) } // TestFSM_RecoveredContinue verifies the error → PrepareStep recovery path. @@ -146,7 +147,7 @@ func TestFSM_RecoveredContinue(t *testing.T) { transitions := obs.Snapshot() last := transitions[len(transitions)-1] - assert.Equal(t, ReActStatePrepareStep, last.To) + assert.Empty(t, cmp.Diff(ReActStatePrepareStep, last.To)) } // TestFSM_UnhandledTriggerIsPermissive verifies that firing an invalid trigger @@ -197,8 +198,8 @@ func TestFSM_StepIndex(t *testing.T) { transitions := obs.Snapshot() require.Len(t, transitions, 2) - assert.Equal(t, 0, transitions[0].StepIndex) - assert.Equal(t, 1, transitions[1].StepIndex) + assert.Empty(t, cmp.Diff(0, transitions[0].StepIndex)) + assert.Empty(t, cmp.Diff(1, transitions[1].StepIndex)) } // TestFSM_NilObserverSafe verifies no panic when no observer is attached. diff --git a/internal/fantasy/schema/schema_test.go b/internal/fantasy/schema/schema_test.go index fdfbfb8c5..770c7419b 100644 --- a/internal/fantasy/schema/schema_test.go +++ b/internal/fantasy/schema/schema_test.go @@ -4,6 +4,8 @@ import ( "reflect" "testing" + "github.com/google/go-cmp/cmp" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -19,8 +21,7 @@ func TestEnumSupport(t *testing.T) { } schema := Generate(reflect.TypeFor[WeatherInput]()) - - require.Equal(t, "object", schema.Type) + assert.Empty(t, cmp.Diff("object", schema.Type)) // Check units field has enum values unitsSchema := schema.Properties["units"] @@ -28,7 +29,7 @@ func TestEnumSupport(t *testing.T) { require.Len(t, unitsSchema.Enum, 3) expectedUnits := []string{"celsius", "fahrenheit", "kelvin"} for i, expected := range expectedUnits { - require.Equal(t, expected, unitsSchema.Enum[i]) + assert.Empty(t, cmp.Diff(expected, unitsSchema.Enum[i])) } // Check required fields (format should not be required due to omitempty) @@ -69,20 +70,20 @@ func TestSchemaToParameters(t *testing.T) { // Check name parameter nameParam, ok := params["name"].(map[string]any) require.True(t, ok, "Expected name parameter to exist") - require.Equal(t, "string", nameParam["type"]) - require.Equal(t, "The name field", nameParam["description"]) + assert.Empty(t, cmp.Diff("string", nameParam["type"])) + assert.Empty(t, cmp.Diff("The name field", nameParam["description"])) // Check age parameter with min/max ageParam, ok := params["age"].(map[string]any) require.True(t, ok, "Expected age parameter to exist") - require.Equal(t, "integer", ageParam["type"]) - require.Equal(t, 0.0, ageParam["minimum"]) - require.Equal(t, 120.0, ageParam["maximum"]) + assert.Empty(t, cmp.Diff("integer", ageParam["type"])) + assert.Empty(t, cmp.Diff(0.0, ageParam["minimum"])) + assert.Empty(t, cmp.Diff(120.0, ageParam["maximum"])) // Check priority parameter with enum priorityParam, ok := params["priority"].(map[string]any) require.True(t, ok, "Expected priority parameter to exist") - require.Equal(t, "string", priorityParam["type"]) + assert.Empty(t, cmp.Diff("string", priorityParam["type"])) enumValues, ok := priorityParam["enum"].([]any) require.True(t, ok) require.Len(t, enumValues, 3) @@ -137,7 +138,7 @@ func TestGenerateSchemaBasicTypes(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() schema := Generate(reflect.TypeOf(tt.input)) - require.Equal(t, tt.expected.Type, schema.Type) + assert.Empty(t, cmp.Diff(tt.expected.Type, schema.Type)) }) } } @@ -180,9 +181,9 @@ func TestGenerateSchemaArrayTypes(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() schema := Generate(reflect.TypeOf(tt.input)) - require.Equal(t, tt.expected.Type, schema.Type) + assert.Empty(t, cmp.Diff(tt.expected.Type, schema.Type)) require.NotNil(t, schema.Items, "Expected items schema to exist") - require.Equal(t, tt.expected.Items.Type, schema.Items.Type) + assert.Empty(t, cmp.Diff(tt.expected.Items.Type, schema.Items.Type)) }) } } @@ -216,7 +217,7 @@ func TestGenerateSchemaMapTypes(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() schema := Generate(reflect.TypeOf(tt.input)) - require.Equal(t, tt.expected, schema.Type) + assert.Empty(t, cmp.Diff(tt.expected, schema.Type)) }) } } @@ -253,10 +254,10 @@ func TestGenerateSchemaStructTypes(t *testing.T) { name: "simple struct", input: SimpleStruct{}, validate: func(t *testing.T, schema Schema) { - require.Equal(t, "object", schema.Type) + assert.Empty(t, cmp.Diff("object", schema.Type)) require.Len(t, schema.Properties, 2) require.NotNil(t, schema.Properties["name"], "Expected name property to exist") - require.Equal(t, "The name field", schema.Properties["name"].Description) + assert.Empty(t, cmp.Diff("The name field", schema.Properties["name"].Description)) require.Len(t, schema.Required, 2) }, }, @@ -265,7 +266,7 @@ func TestGenerateSchemaStructTypes(t *testing.T) { input: StructWithOmitEmpty{}, validate: func(t *testing.T, schema Schema) { require.Len(t, schema.Required, 1) - require.Equal(t, "required", schema.Required[0]) + assert.Empty(t, cmp.Diff("required", schema.Required[0])) }, }, { @@ -305,14 +306,13 @@ func TestGenerateSchemaPointerTypes(t *testing.T) { } schema := Generate(reflect.TypeFor[StructWithPointers]()) - - require.Equal(t, "object", schema.Type) + assert.Empty(t, cmp.Diff("object", schema.Type)) require.NotNil(t, schema.Properties["name"], "Expected name property to exist") - require.Equal(t, "string", schema.Properties["name"].Type) + assert.Empty(t, cmp.Diff("string", schema.Properties["name"].Type)) require.NotNil(t, schema.Properties["age"], "Expected age property to exist") - require.Equal(t, "integer", schema.Properties["age"].Type) + assert.Empty(t, cmp.Diff("integer", schema.Properties["age"].Type)) } func TestGenerateSchemaNestedStructs(t *testing.T) { @@ -329,13 +329,12 @@ func TestGenerateSchemaNestedStructs(t *testing.T) { } schema := Generate(reflect.TypeFor[Person]()) - - require.Equal(t, "object", schema.Type) + assert.Empty(t, cmp.Diff("object", schema.Type)) require.NotNil(t, schema.Properties["address"], "Expected address property to exist") addressSchema := schema.Properties["address"] - require.Equal(t, "object", addressSchema.Type) + assert.Empty(t, cmp.Diff("object", addressSchema.Type)) require.NotNil(t, addressSchema.Properties["street"], "Expected street property in address to exist") require.NotNil(t, addressSchema.Properties["city"], "Expected city property in address to exist") @@ -350,8 +349,7 @@ func TestGenerateSchemaRecursiveStructs(t *testing.T) { } schema := Generate(reflect.TypeFor[Node]()) - - require.Equal(t, "object", schema.Type) + assert.Empty(t, cmp.Diff("object", schema.Type)) require.NotNil(t, schema.Properties["value"], "Expected value property to exist") @@ -359,7 +357,7 @@ func TestGenerateSchemaRecursiveStructs(t *testing.T) { // The recursive reference should be handled gracefully nextSchema := schema.Properties["next"] - require.Equal(t, "object", nextSchema.Type) + assert.Empty(t, cmp.Diff("object", nextSchema.Type)) } func TestGenerateSchemaWithEnumTags(t *testing.T) { @@ -379,7 +377,7 @@ func TestGenerateSchemaWithEnumTags(t *testing.T) { require.Len(t, levelSchema.Enum, 4) expectedLevels := []string{"debug", "info", "warn", "error"} for i, expected := range expectedLevels { - require.Equal(t, expected, levelSchema.Enum[i]) + assert.Empty(t, cmp.Diff(expected, levelSchema.Enum[i])) } // Check format field @@ -407,24 +405,24 @@ func TestGenerateSchemaComplexTypes(t *testing.T) { // Check string slice stringSliceSchema := schema.Properties["string_slice"] require.NotNil(t, stringSliceSchema, "Expected string_slice property to exist") - require.Equal(t, "array", stringSliceSchema.Type) - require.Equal(t, "string", stringSliceSchema.Items.Type) + assert.Empty(t, cmp.Diff("array", stringSliceSchema.Type)) + assert.Empty(t, cmp.Diff("string", stringSliceSchema.Items.Type)) // Check int map intMapSchema := schema.Properties["int_map"] require.NotNil(t, intMapSchema, "Expected int_map property to exist") - require.Equal(t, "object", intMapSchema.Type) + assert.Empty(t, cmp.Diff("object", intMapSchema.Type)) // Check nested slice nestedSliceSchema := schema.Properties["nested_slice"] require.NotNil(t, nestedSliceSchema, "Expected nested_slice property to exist") - require.Equal(t, "array", nestedSliceSchema.Type) - require.Equal(t, "object", nestedSliceSchema.Items.Type) + assert.Empty(t, cmp.Diff("array", nestedSliceSchema.Type)) + assert.Empty(t, cmp.Diff("object", nestedSliceSchema.Items.Type)) // Check interface interfaceSchema := schema.Properties["interface"] require.NotNil(t, interfaceSchema, "Expected interface property to exist") - require.Equal(t, "object", interfaceSchema.Type) + assert.Empty(t, cmp.Diff("object", interfaceSchema.Type)) } func TestToSnakeCase(t *testing.T) { @@ -449,7 +447,7 @@ func TestToSnakeCase(t *testing.T) { t.Run(tt.input, func(t *testing.T) { t.Parallel() result := toSnakeCase(tt.input) - require.Equal(t, tt.expected, result, "toSnakeCase(%s)", tt.input) + assert.Empty(t, cmp.Diff(tt.expected, result), "toSnakeCase(%s)", tt.input) }) } } @@ -530,7 +528,7 @@ func TestSchemaToParametersEdgeCases(t *testing.T) { resultParam := result[key].(map[string]any) expectedParam := expectedValue.(map[string]any) for propKey, propValue := range expectedParam { - require.Equal(t, propValue, resultParam[propKey], "Expected %s.%s", key, propKey) + assert.Empty(t, cmp.Diff(propValue, resultParam[propKey]), "Expected %s.%s", key, propKey) } } }) @@ -564,7 +562,7 @@ func TestNormalize_TypeArray(t *testing.T) { require.Contains(t, variant, "items") } } - require.Equal(t, "Config value", val["description"]) + assert.Empty(t, cmp.Diff("Config value", val["description"])) } func TestNormalize_SingleStringType(t *testing.T) { @@ -580,7 +578,7 @@ func TestNormalize_SingleStringType(t *testing.T) { Normalize(node) val := node["properties"].(map[string]any)["name"].(map[string]any) - require.Equal(t, "string", val["type"]) + assert.Empty(t, cmp.Diff("string", val["type"])) } func TestNormalize_BareArrayGetsItems(t *testing.T) { @@ -596,7 +594,7 @@ func TestNormalize_BareArrayGetsItems(t *testing.T) { Normalize(node) val := node["properties"].(map[string]any)["tags"].(map[string]any) - require.Equal(t, "array", val["type"]) + assert.Empty(t, cmp.Diff("array", val["type"])) require.Contains(t, val, "items") } @@ -617,7 +615,7 @@ func TestNormalize_SingleElementTypeArray(t *testing.T) { anyOf, ok := val["anyOf"].([]any) require.True(t, ok) require.Len(t, anyOf, 1) - require.Equal(t, "string", anyOf[0].(map[string]any)["type"]) + assert.Empty(t, cmp.Diff("string", anyOf[0].(map[string]any)["type"])) } func TestNormalize_NestedProperties(t *testing.T) { diff --git a/internal/fantasy/tool_runtime_dag_test.go b/internal/fantasy/tool_runtime_dag_test.go index c5e15d090..c0deacec7 100644 --- a/internal/fantasy/tool_runtime_dag_test.go +++ b/internal/fantasy/tool_runtime_dag_test.go @@ -8,6 +8,8 @@ import ( "testing" "time" + "github.com/google/go-cmp/cmp" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -119,10 +121,11 @@ func TestDAGToolRuntime_DependenciesWaitAndInputIsResolved(t *testing.T) { require.NoError(t, err) require.Len(t, res, 2) + assert. - // B should have received val=1 and returned "1". - require.Equal(t, "callB", res[1].ToolCallID) - require.Equal(t, "1", res[1].Result.(ToolResultOutputContentText).Text) + // B should have received val=1 and returned "1". + Empty(t, cmp.Diff("callB", res[1].ToolCallID)) + assert.Empty(t, cmp.Diff("1", res[1].Result.(ToolResultOutputContentText).Text)) } func TestDAGToolRuntime_CycleDetected(t *testing.T) { @@ -178,7 +181,7 @@ func TestDAGToolRuntime_OnToolResultSerialized(t *testing.T) { orderMu.Lock() defer orderMu.Unlock() - require.Equal(t, []string{"a", "b"}, order) + assert.Empty(t, cmp.Diff([]string{"a", "b"}, order)) } func TestDAGToolRuntime_MetricsAndLogHooks(t *testing.T) { diff --git a/internal/fantasy/tool_runtime_parallel_test.go b/internal/fantasy/tool_runtime_parallel_test.go index 7e5c429ce..0a64f9367 100644 --- a/internal/fantasy/tool_runtime_parallel_test.go +++ b/internal/fantasy/tool_runtime_parallel_test.go @@ -7,6 +7,8 @@ import ( "testing" "time" + "github.com/google/go-cmp/cmp" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -49,16 +51,15 @@ func TestParallelToolRuntime_OrderAndCallbackDeterminism(t *testing.T) { results, err := runtime.Execute(t.Context(), []AgentTool{tool}, toolCalls, cb) require.NoError(t, err) require.Len(t, results, 3) - - require.Equal(t, "c1", results[0].ToolCallID) - require.Equal(t, "a", results[0].Result.(ToolResultOutputContentText).Text) - require.Equal(t, "c2", results[1].ToolCallID) - require.Equal(t, "b", results[1].Result.(ToolResultOutputContentText).Text) - require.Equal(t, "c3", results[2].ToolCallID) - require.Equal(t, "c", results[2].Result.(ToolResultOutputContentText).Text) + assert.Empty(t, cmp.Diff("c1", results[0].ToolCallID)) + assert.Empty(t, cmp.Diff("a", results[0].Result.(ToolResultOutputContentText).Text)) + assert.Empty(t, cmp.Diff("c2", results[1].ToolCallID)) + assert.Empty(t, cmp.Diff("b", results[1].Result.(ToolResultOutputContentText).Text)) + assert.Empty(t, cmp.Diff("c3", results[2].ToolCallID)) + assert.Empty(t, cmp.Diff("c", results[2].Result.(ToolResultOutputContentText).Text)) cbMu.Lock() - require.Equal(t, []string{"c1", "c2", "c3"}, cbOrder) + assert.Empty(t, cmp.Diff([]string{"c1", "c2", "c3"}, cbOrder)) cbMu.Unlock() } @@ -97,7 +98,7 @@ func TestParallelToolRuntime_BarrierForNonParallelTools(t *testing.T) { require.Len(t, results, 4) mu.Lock() - require.Equal(t, []string{"p1", "p2", "s1", "p3"}, order) + assert.Empty(t, cmp.Diff([]string{"p1", "p2", "s1", "p3"}, order)) mu.Unlock() } diff --git a/internal/fantasy/tool_test.go b/internal/fantasy/tool_test.go index 4763c462a..abb4bf784 100644 --- a/internal/fantasy/tool_test.go +++ b/internal/fantasy/tool_test.go @@ -6,6 +6,7 @@ import ( "testing" "github.com/google/go-cmp/cmp" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -30,9 +31,9 @@ func TestTypedToolFuncExample(t *testing.T) { // Check the tool info info := tool.Info() - require.Equal(t, "calculator", info.Name) + assert.Empty(t, cmp.Diff("calculator", info.Name)) require.Len(t, info.Required, 1) - require.Equal(t, "expression", info.Required[0]) + assert.Empty(t, cmp.Diff("expression", info.Required[0])) // Test execution call := ToolCall{ @@ -43,7 +44,7 @@ func TestTypedToolFuncExample(t *testing.T) { result, err := tool.Run(t.Context(), call) require.NoError(t, err) - require.Equal(t, "4", result.Content) + assert.Empty(t, cmp.Diff("4", result.Content)) require.False(t, result.IsError) } diff --git a/pkg/agent/kv_delegate_test.go b/pkg/agent/kv_delegate_test.go index 4af889f1a..4f2f9cc01 100644 --- a/pkg/agent/kv_delegate_test.go +++ b/pkg/agent/kv_delegate_test.go @@ -3,6 +3,7 @@ package agent_test import ( "testing" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -24,7 +25,7 @@ func TestDelegateKV_PutAndGet(t *testing.T) { got, err := kv.Get(ctx, "key1") require.NoError(t, err) - assert.Equal(t, []byte("hello world"), got) + assert.Empty(t, cmp.Diff([]byte("hello world"), got)) } func TestDelegateKV_GetMissingKey(t *testing.T) { @@ -47,7 +48,7 @@ func TestDelegateKV_PutOverwrite(t *testing.T) { got, err := kv.Get(ctx, "k") require.NoError(t, err) - assert.Equal(t, []byte("v2"), got, "second put should overwrite the first") + assert.Empty(t, cmp.Diff([]byte("v2"), got), "second put should overwrite the first") } func TestDelegateKV_BinaryValues(t *testing.T) { @@ -61,7 +62,7 @@ func TestDelegateKV_BinaryValues(t *testing.T) { got, err := kv.Get(ctx, "binary-key") require.NoError(t, err) - assert.Equal(t, data, got, "binary round-trip must be lossless") + assert.Empty(t, cmp.Diff(data, got), "binary round-trip must be lossless") } func TestDelegateKV_Scan(t *testing.T) { @@ -101,7 +102,7 @@ func TestDelegateKV_ScanSorted(t *testing.T) { got, err := kv.Scan(ctx, "z/") require.NoError(t, err) - assert.Equal(t, []string{"z/a", "z/b", "z/c"}, got, "Scan must return keys sorted") + assert.Empty(t, cmp.Diff([]string{"z/a", "z/b", "z/c"}, got), "Scan must return keys sorted") } func TestDelegateKV_AgentIsolation(t *testing.T) { @@ -116,11 +117,11 @@ func TestDelegateKV_AgentIsolation(t *testing.T) { v1, err := kv1.Get(ctx, "shared-key") require.NoError(t, err) - assert.Equal(t, []byte("from-A"), v1) + assert.Empty(t, cmp.Diff([]byte("from-A"), v1)) v2, err := kv2.Get(ctx, "shared-key") require.NoError(t, err) - assert.Equal(t, []byte("from-B"), v2) + assert.Empty(t, cmp.Diff([]byte("from-B"), v2)) } func TestDelegateKV_EmptyKey_PutErrors(t *testing.T) { diff --git a/pkg/agent/offloading_runtime_test.go b/pkg/agent/offloading_runtime_test.go index fea02828e..76277eba0 100644 --- a/pkg/agent/offloading_runtime_test.go +++ b/pkg/agent/offloading_runtime_test.go @@ -6,6 +6,7 @@ import ( "testing" "charm.land/fantasy" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -85,7 +86,7 @@ func TestOffloading_SmallResult_KeptInline(t *testing.T) { require.NoError(t, err) require.Len(t, results, 1) - assert.Equal(t, "echo:call-a", results[0].Result.(fantasy.ToolResultOutputContentText).Text) + assert.Empty(t, cmp.Diff("echo:call-a", results[0].Result.(fantasy.ToolResultOutputContentText).Text)) // KV should still have the full result stored keys, err := kv.Scan(ctx, "tool_results/") @@ -146,7 +147,7 @@ func TestOffloading_LargeResult_Truncated(t *testing.T) { require.NoError(t, err) full.Write(part) } - assert.Equal(t, longText, full.String(), "reassembled chunks must equal original text") + assert.Empty(t, cmp.Diff(longText, full.String()), "reassembled chunks must equal original text") } // TestOffloading_DBMetadata_Inserted verifies that the DB record is created. @@ -290,7 +291,7 @@ func TestChunkString(t *testing.T) { } if tt.chunkSize > 0 && len(tt.input) > effectiveThreshold { - assert.Equal(t, tt.wantLen, chunkCount) + assert.Empty(t, cmp.Diff(tt.wantLen, chunkCount)) } }) } diff --git a/pkg/agent/state_store_test.go b/pkg/agent/state_store_test.go index aa5967d2c..782c0ad62 100644 --- a/pkg/agent/state_store_test.go +++ b/pkg/agent/state_store_test.go @@ -204,12 +204,12 @@ func TestCheckpointStore_CreateAndList(t *testing.T) { require.NoError(t, err) assert.False(t, cp.ID.IsZero()) - assert.Equal(t, "after-step-0", cp.Name) + assert.Empty(t, cmp.Diff("after-step-0", cp.Name)) cps, err := cs.ListCheckpoints(ctx, convID) require.NoError(t, err) assert.Len(t, cps, 1) - assert.Equal(t, cp.ID, cps[0].ID) + assert.Empty(t, cmp.Diff(cp.ID, cps[0].ID)) } func TestCheckpointStore_GetByName(t *testing.T) { @@ -232,7 +232,7 @@ func TestCheckpointStore_GetByName(t *testing.T) { cp, err := cs.GetCheckpoint(ctx, convID, "snap-1") require.NoError(t, err) - assert.Equal(t, "snap-1", cp.Name) + assert.Empty(t, cmp.Diff("snap-1", cp.Name)) } func TestCheckpointStore_EmptyName(t *testing.T) { diff --git a/pkg/agent/tool_result_search_test.go b/pkg/agent/tool_result_search_test.go index 703604657..bbf28007f 100644 --- a/pkg/agent/tool_result_search_test.go +++ b/pkg/agent/tool_result_search_test.go @@ -7,6 +7,7 @@ import ( jsonv2 "github.com/go-json-experiment/json" "charm.land/fantasy" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -149,12 +150,12 @@ func TestToolResultSearch_QueryFilters(t *testing.T) { out := invokeSearch(t, tool, input) total, _ := out["total"].(float64) - assert.Equal(t, tt.wantTotal, total) + assert.Empty(t, cmp.Diff(tt.wantTotal, total)) if tt.wantTCID != "" { items := out["items"].([]any) item := items[0].(map[string]any) - assert.Equal(t, tt.wantTCID, item["tool_call_id"]) + assert.Empty(t, cmp.Diff(tt.wantTCID, item["tool_call_id"])) } }) } @@ -229,7 +230,7 @@ func TestToolResultSearch_LineView(t *testing.T) { // The view_range reflects the actual clamped range (1 JSON line available). viewRange := item["view_range"].(map[string]any) - assert.Equal(t, float64(1), viewRange["start_line"], "clamped to 1 since full.json is a single-line JSON blob") + assert.Empty(t, cmp.Diff(float64(1), viewRange["start_line"]), "clamped to 1 since full.json is a single-line JSON blob") } // TestToolResultSearch_ChunkView verifies chunk-range retrieval from KV. @@ -282,5 +283,5 @@ func TestToolResultSearch_ChunkView(t *testing.T) { view := item["view"].(string) // First two chunks: "abcdefghijabcdefghij" + "abcdefghijabcdefghij" = 40 chars - assert.Equal(t, strings.Repeat("abcdefghij", 4), view, "chunk 0+1 should be first 40 chars") + assert.Empty(t, cmp.Diff(strings.Repeat("abcdefghij", 4), view), "chunk 0+1 should be first 40 chars") } diff --git a/pkg/itr/commands_test.go b/pkg/itr/commands_test.go index 219dc4625..d9e94ccc5 100644 --- a/pkg/itr/commands_test.go +++ b/pkg/itr/commands_test.go @@ -87,8 +87,8 @@ func TestToolExecPayloadPreservation(t *testing.T) { payload, ok := decoded.Payload.(ToolExec) require.True(t, ok, "payload should be ToolExec") - assert.Equal(t, "shell", payload.ToolName) - assert.Equal(t, `{"cmd":"ls -la"}`, payload.ArgsJSON) + assert.Empty(t, cmp.Diff("shell", payload.ToolName)) + assert.Empty(t, cmp.Diff(`{"cmd":"ls -la"}`, payload.ArgsJSON)) } func TestGrepPayloadPreservation(t *testing.T) { @@ -102,8 +102,8 @@ func TestGrepPayloadPreservation(t *testing.T) { payload, ok := decoded.Payload.(Grep) require.True(t, ok) - assert.Equal(t, "error.*fatal", payload.Pattern) - assert.Equal(t, uint32(25), payload.MaxMatches) + assert.Empty(t, cmp.Diff("error.*fatal", payload.Pattern)) + assert.Empty(t, cmp.Diff(uint32(25), payload.MaxMatches)) assert.True(t, payload.CaseInsensitive) } @@ -127,20 +127,20 @@ func TestDAGPlanPayloadPreservation(t *testing.T) { dagPlan, ok := decoded.Payload.(DAGPlan) require.True(t, ok) assert.Len(t, dagPlan.Nodes, 2) - assert.Equal(t, "a", dagPlan.Nodes[0].ID) - assert.Equal(t, CmdToolSearch, dagPlan.Nodes[0].Type) - assert.Equal(t, []string{"a"}, dagPlan.Nodes[1].DependsOn) - assert.Equal(t, uint8(2), dagPlan.MaxParallel) - assert.Equal(t, "combine results", dagPlan.JoinerQuery) + assert.Empty(t, cmp.Diff("a", dagPlan.Nodes[0].ID)) + assert.Empty(t, cmp.Diff(CmdToolSearch, dagPlan.Nodes[0].Type)) + assert.Empty(t, cmp.Diff([]string{"a"}, dagPlan.Nodes[1].DependsOn)) + assert.Empty(t, cmp.Diff(uint8(2), dagPlan.MaxParallel)) + assert.Empty(t, cmp.Diff("combine results", dagPlan.JoinerQuery)) ts, ok := dagPlan.Nodes[0].Payload.(ToolSearch) require.True(t, ok, "ToolSearch payload should be retyped after unmarshal") - assert.Equal(t, "files", ts.Query) - assert.Equal(t, uint8(5), ts.MaxResults) + assert.Empty(t, cmp.Diff("files", ts.Query)) + assert.Empty(t, cmp.Diff(uint8(5), ts.MaxResults)) te, ok := dagPlan.Nodes[1].Payload.(ToolExec) require.True(t, ok, "ToolExec payload should be retyped after unmarshal") - assert.Equal(t, "read", te.ToolName) + assert.Empty(t, cmp.Diff("read", te.ToolName)) assert.Contains(t, te.ArgsJSON, "#nodea") } @@ -214,9 +214,9 @@ func TestRequestJSON_Roundtrip(t *testing.T) { decoded, err := UnmarshalRequestJSON(data) require.NoError(t, err) - assert.Equal(t, orig.ID, decoded.ID) + assert.Empty(t, cmp.Diff(orig.ID, decoded.ID)) p := decoded.Payload.(ToolExec) - assert.Equal(t, "shell", p.ToolName) + assert.Empty(t, cmp.Diff("shell", p.ToolName)) } func TestResponseJSON_Roundtrip(t *testing.T) { @@ -227,6 +227,6 @@ func TestResponseJSON_Roundtrip(t *testing.T) { decoded, err := UnmarshalResponseJSON(data) require.NoError(t, err) - assert.Equal(t, orig.ID, decoded.ID) - assert.Equal(t, "ok", decoded.Result) + assert.Empty(t, cmp.Diff(orig.ID, decoded.ID)) + assert.Empty(t, cmp.Diff("ok", decoded.Result)) } diff --git a/pkg/itr/dag/executor_test.go b/pkg/itr/dag/executor_test.go index 3020515fb..eb61e8b1b 100644 --- a/pkg/itr/dag/executor_test.go +++ b/pkg/itr/dag/executor_test.go @@ -11,6 +11,7 @@ import ( "github.com/ZanzyTHEbar/dragonscale/pkg/itr/dag" "github.com/ZanzyTHEbar/dragonscale/pkg/security/securebus" "github.com/ZanzyTHEbar/dragonscale/pkg/tools" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -95,8 +96,8 @@ func TestExecutor_ParallelNodes(t *testing.T) { result, err := executor.Execute(t.Context(), "test-sess", plan) require.NoError(t, err) - assert.Equal(t, "a-result", result.NodeResults["n1"]) - assert.Equal(t, "b-result", result.NodeResults["n2"]) + assert.Empty(t, cmp.Diff("a-result", result.NodeResults["n1"])) + assert.Empty(t, cmp.Diff("b-result", result.NodeResults["n2"])) } func TestExecutor_CycleDetection(t *testing.T) { @@ -148,7 +149,7 @@ func TestExecutor_WithJoiner(t *testing.T) { result, err := executor.Execute(t.Context(), "test-sess", plan) require.NoError(t, err) assert.Contains(t, result.FinalAnswer, "synthesized:") - assert.Equal(t, uint32(50), result.TotalTokens) + assert.Empty(t, cmp.Diff(uint32(50), result.TotalTokens)) } func TestExecutor_EmptyPlan(t *testing.T) { @@ -192,21 +193,21 @@ func TestRouter_SimpleQuerySelectsReAct(t *testing.T) { t.Parallel() cfg := dag.DefaultRouterConfig() mode := dag.Route(dag.ModeAuto, "What is the weather?", cfg) - assert.Equal(t, dag.ModeReAct, mode) + assert.Empty(t, cmp.Diff(dag.ModeReAct, mode)) } func TestRouter_ComplexQuerySelectsDAG(t *testing.T) { t.Parallel() cfg := dag.DefaultRouterConfig() mode := dag.Route(dag.ModeAuto, "Search for the latest news about AI, read the top 3 articles, and compare their viewpoints to create a summary report with aggregate statistics", cfg) - assert.Equal(t, dag.ModeDAG, mode) + assert.Empty(t, cmp.Diff(dag.ModeDAG, mode)) } func TestRouter_ExplicitModeOverridesAuto(t *testing.T) { t.Parallel() cfg := dag.DefaultRouterConfig() mode := dag.Route(dag.ModeReAct, "Do many complex parallel things simultaneously", cfg) - assert.Equal(t, dag.ModeReAct, mode) + assert.Empty(t, cmp.Diff(dag.ModeReAct, mode)) } func TestPlanner_ValidatePlan(t *testing.T) { diff --git a/pkg/itr/dag/planner_test.go b/pkg/itr/dag/planner_test.go index 7fe7ba5f3..b14c7015b 100644 --- a/pkg/itr/dag/planner_test.go +++ b/pkg/itr/dag/planner_test.go @@ -7,6 +7,7 @@ import ( jsonv2 "github.com/go-json-experiment/json" "github.com/ZanzyTHEbar/dragonscale/pkg/itr" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -14,34 +15,34 @@ import ( func TestExtractJSON_PlainJSON(t *testing.T) { t.Parallel() input := `{"nodes": [{"id": "n1"}]}` - assert.Equal(t, input, extractJSON(input)) + assert.Empty(t, cmp.Diff(input, extractJSON(input))) } func TestExtractJSON_MarkdownFenced(t *testing.T) { t.Parallel() input := "Here is the plan:\n```json\n{\"nodes\": [{\"id\": \"n1\"}]}\n```\nDone." - assert.Equal(t, `{"nodes": [{"id": "n1"}]}`, extractJSON(input)) + assert.Empty(t, cmp.Diff(`{"nodes": [{"id": "n1"}]}`, extractJSON(input))) } func TestExtractJSON_GenericFenced(t *testing.T) { t.Parallel() input := "```\n{\"nodes\": []}\n```" - assert.Equal(t, `{"nodes": []}`, extractJSON(input)) + assert.Empty(t, cmp.Diff(`{"nodes": []}`, extractJSON(input))) } func TestExtractJSON_LeadingText(t *testing.T) { t.Parallel() input := "The plan is: {\"nodes\":[]}" - assert.Equal(t, `{"nodes":[]}`, extractJSON(input)) + assert.Empty(t, cmp.Diff(`{"nodes":[]}`, extractJSON(input))) } func TestFindIndex(t *testing.T) { t.Parallel() - assert.Equal(t, 0, findIndex("abc", "a")) - assert.Equal(t, 2, findIndex("abc", "c")) - assert.Equal(t, -1, findIndex("abc", "z")) - assert.Equal(t, -1, findIndex("", "a")) - assert.Equal(t, -1, findIndex("ab", "abc")) + assert.Empty(t, cmp.Diff(0, findIndex("abc", "a"))) + assert.Empty(t, cmp.Diff(2, findIndex("abc", "c"))) + assert.Empty(t, cmp.Diff(-1, findIndex("abc", "z"))) + assert.Empty(t, cmp.Diff(-1, findIndex("", "a"))) + assert.Empty(t, cmp.Diff(-1, findIndex("ab", "abc"))) } func TestValidatePlan_Valid(t *testing.T) { @@ -135,13 +136,13 @@ func TestParsePlanResponse_ValidJSON(t *testing.T) { plan, err := parsePlanResponse(input) require.NoError(t, err) require.Len(t, plan.Nodes, 1) - assert.Equal(t, "n1", plan.Nodes[0].ID) - assert.Equal(t, itr.CmdToolExec, plan.Nodes[0].Type) - assert.Equal(t, "summarize", plan.JoinerQuery) + assert.Empty(t, cmp.Diff("n1", plan.Nodes[0].ID)) + assert.Empty(t, cmp.Diff(itr.CmdToolExec, plan.Nodes[0].Type)) + assert.Empty(t, cmp.Diff("summarize", plan.JoinerQuery)) te, ok := plan.Nodes[0].Payload.(itr.ToolExec) require.True(t, ok) - assert.Equal(t, "read_file", te.ToolName) + assert.Empty(t, cmp.Diff("read_file", te.ToolName)) } func TestParsePlanResponse_WithMarkdownFence(t *testing.T) { @@ -154,7 +155,7 @@ func TestParsePlanResponse_WithMarkdownFence(t *testing.T) { ts, ok := plan.Nodes[0].Payload.(itr.ToolSearch) require.True(t, ok) - assert.Equal(t, "files", ts.Query) + assert.Empty(t, cmp.Diff("files", ts.Query)) } func TestParsePlanResponse_InvalidJSON(t *testing.T) { @@ -183,10 +184,10 @@ func TestPlannerPlanE2E(t *testing.T) { planner := NewPlanner(mockLLM, nil, DefaultPlannerConfig()) plan, tokens, err := planner.Plan(t.Context(), "search and read", nil) require.NoError(t, err) - assert.Equal(t, uint32(100), tokens) + assert.Empty(t, cmp.Diff(uint32(100), tokens)) require.Len(t, plan.Nodes, 2) - assert.Equal(t, "search", plan.Nodes[0].ID) - assert.Equal(t, uint8(8), plan.MaxParallel) + assert.Empty(t, cmp.Diff("search", plan.Nodes[0].ID)) + assert.Empty(t, cmp.Diff(uint8(8), plan.MaxParallel)) } func TestPlannerPlanLLMError(t *testing.T) { @@ -198,5 +199,5 @@ func TestPlannerPlanLLMError(t *testing.T) { planner := NewPlanner(mockLLM, nil, DefaultPlannerConfig()) _, tokens, err := planner.Plan(t.Context(), "anything", nil) assert.Error(t, err) - assert.Equal(t, uint32(50), tokens) + assert.Empty(t, cmp.Diff(uint32(50), tokens)) } diff --git a/pkg/itr/dag/replan_test.go b/pkg/itr/dag/replan_test.go index a75568f9d..0db012750 100644 --- a/pkg/itr/dag/replan_test.go +++ b/pkg/itr/dag/replan_test.go @@ -3,6 +3,7 @@ package dag import ( "testing" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" ) @@ -16,5 +17,5 @@ func TestNeedsReplan(t *testing.T) { func TestReplanSentinelIsConsistent(t *testing.T) { t.Parallel() - assert.Equal(t, "[NEEDS_MORE_STEPS]", replanSentinel) + assert.Empty(t, cmp.Diff("[NEEDS_MORE_STEPS]", replanSentinel)) } diff --git a/pkg/itr/dag/resolver_test.go b/pkg/itr/dag/resolver_test.go index 4a0ed5cec..c78a3ea31 100644 --- a/pkg/itr/dag/resolver_test.go +++ b/pkg/itr/dag/resolver_test.go @@ -3,6 +3,7 @@ package dag import ( "testing" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -18,9 +19,9 @@ func TestTopologicalOrderLinear(t *testing.T) { waves, err := topologicalOrder(states) require.NoError(t, err) require.Len(t, waves, 3) - assert.Equal(t, []string{"a"}, waves[0]) - assert.Equal(t, []string{"b"}, waves[1]) - assert.Equal(t, []string{"c"}, waves[2]) + assert.Empty(t, cmp.Diff([]string{"a"}, waves[0])) + assert.Empty(t, cmp.Diff([]string{"b"}, waves[1])) + assert.Empty(t, cmp.Diff([]string{"c"}, waves[2])) } func TestTopologicalOrderParallel(t *testing.T) { @@ -38,7 +39,7 @@ func TestTopologicalOrderParallel(t *testing.T) { assert.Len(t, waves[0], 2) assert.Contains(t, waves[0], "a") assert.Contains(t, waves[0], "b") - assert.Equal(t, []string{"c"}, waves[1]) + assert.Empty(t, cmp.Diff([]string{"c"}, waves[1])) } func TestTopologicalOrderCycleDetection(t *testing.T) { @@ -63,7 +64,7 @@ func TestTopologicalOrderSingleNode(t *testing.T) { waves, err := topologicalOrder(states) require.NoError(t, err) require.Len(t, waves, 1) - assert.Equal(t, []string{"only"}, waves[0]) + assert.Empty(t, cmp.Diff([]string{"only"}, waves[0])) } func TestResolveRefs(t *testing.T) { @@ -84,7 +85,7 @@ func TestResolveRefsNoMatch(t *testing.T) { states := map[string]*nodeState{} input := `{"path":"#nodemissing"}` result := resolveRefs(input, states) - assert.Equal(t, input, result) + assert.Empty(t, cmp.Diff(input, result)) } func TestResolveToolExecArgsNoRefs(t *testing.T) { @@ -92,7 +93,7 @@ func TestResolveToolExecArgsNoRefs(t *testing.T) { states := map[string]*nodeState{} input := `{"path":"/tmp/plain.txt"}` result := resolveToolExecArgs(input, states) - assert.Equal(t, input, result) + assert.Empty(t, cmp.Diff(input, result)) } func TestEscapeForJSON(t *testing.T) { @@ -108,7 +109,7 @@ func TestEscapeForJSON(t *testing.T) { } for _, tt := range tests { t.Run(tt.input, func(t *testing.T) { - assert.Equal(t, tt.expected, escapeForJSON(tt.input)) + assert.Empty(t, cmp.Diff(tt.expected, escapeForJSON(tt.input))) }) } } @@ -125,5 +126,5 @@ func TestNodeStateSetAndGetResult(t *testing.T) { result, err := ns.getResult() assert.NoError(t, err) - assert.Equal(t, "result-data", result) + assert.Empty(t, cmp.Diff("result-data", result)) } diff --git a/pkg/itr/dag/router_test.go b/pkg/itr/dag/router_test.go index 0a2224c05..f3ae0dc51 100644 --- a/pkg/itr/dag/router_test.go +++ b/pkg/itr/dag/router_test.go @@ -4,6 +4,7 @@ import ( "strings" "testing" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" ) @@ -11,21 +12,21 @@ func TestRouteExplicitModes(t *testing.T) { t.Parallel() cfg := DefaultRouterConfig() - assert.Equal(t, ModeReAct, Route(ModeReAct, "anything", cfg)) - assert.Equal(t, ModeDAG, Route(ModeDAG, "anything", cfg)) + assert.Empty(t, cmp.Diff(ModeReAct, Route(ModeReAct, "anything", cfg))) + assert.Empty(t, cmp.Diff(ModeDAG, Route(ModeDAG, "anything", cfg))) } func TestRouteAutoSimpleQuery(t *testing.T) { t.Parallel() cfg := DefaultRouterConfig() - assert.Equal(t, ModeReAct, Route(ModeAuto, "what is the weather?", cfg)) + assert.Empty(t, cmp.Diff(ModeReAct, Route(ModeAuto, "what is the weather?", cfg))) } func TestRouteAutoComplexQuery(t *testing.T) { t.Parallel() cfg := DefaultRouterConfig() longQuery := strings.Repeat("word ", 35) - assert.Equal(t, ModeDAG, Route(ModeAuto, longQuery, cfg)) + assert.Empty(t, cmp.Diff(ModeDAG, Route(ModeAuto, longQuery, cfg))) } func TestRouteAutoParallelKeywords(t *testing.T) { @@ -40,7 +41,7 @@ func TestRouteAutoParallelKeywords(t *testing.T) { } for _, q := range keywords { t.Run(q, func(t *testing.T) { - assert.Equal(t, ModeDAG, Route(ModeAuto, q, cfg)) + assert.Empty(t, cmp.Diff(ModeDAG, Route(ModeAuto, q, cfg))) }) } } @@ -49,19 +50,19 @@ func TestRouteAutoToolSignals(t *testing.T) { t.Parallel() cfg := DefaultRouterConfig() q := "search the codebase, read the file, then execute the command" - assert.Equal(t, ModeDAG, Route(ModeAuto, q, cfg)) + assert.Empty(t, cmp.Diff(ModeDAG, Route(ModeAuto, q, cfg))) } func TestToolLoopModeString(t *testing.T) { t.Parallel() - assert.Equal(t, "react", ModeReAct.String()) - assert.Equal(t, "dag", ModeDAG.String()) - assert.Equal(t, "auto", ModeAuto.String()) - assert.Equal(t, "unknown", ToolLoopMode(99).String()) + assert.Empty(t, cmp.Diff("react", ModeReAct.String())) + assert.Empty(t, cmp.Diff("dag", ModeDAG.String())) + assert.Empty(t, cmp.Diff("auto", ModeAuto.String())) + assert.Empty(t, cmp.Diff("unknown", ToolLoopMode(99).String())) } func TestClassifyQueryDefault(t *testing.T) { t.Parallel() cfg := DefaultRouterConfig() - assert.Equal(t, ModeReAct, classifyQuery("hello", cfg)) + assert.Empty(t, cmp.Diff(ModeReAct, classifyQuery("hello", cfg))) } diff --git a/pkg/itr/wasm/runtime_test.go b/pkg/itr/wasm/runtime_test.go index a60d6c9b2..166b7c0b7 100644 --- a/pkg/itr/wasm/runtime_test.go +++ b/pkg/itr/wasm/runtime_test.go @@ -4,6 +4,7 @@ import ( "testing" "time" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -62,7 +63,7 @@ func TestExecuteMinimalModule(t *testing.T) { result, err := rt.Execute(ctx, minimalWASM, "") require.NoError(t, err) - assert.Equal(t, uint32(0), result.ExitCode) + assert.Empty(t, cmp.Diff(uint32(0), result.ExitCode)) assert.Empty(t, result.Stderr) assert.True(t, result.Duration > 0) } @@ -80,7 +81,7 @@ func TestExecuteTimeout(t *testing.T) { // Minimal module is fast enough to succeed even with 1ms timeout. // This test validates that the timeout machinery doesn't break normal execution. if err == nil { - assert.Equal(t, uint32(0), result.ExitCode) + assert.Empty(t, cmp.Diff(uint32(0), result.ExitCode)) } } @@ -123,8 +124,8 @@ func TestLimitedBuffer(t *testing.T) { lb := &limitedBuffer{max: 5} n, err := lb.Write([]byte("hello world")) assert.NoError(t, err) - assert.Equal(t, 11, n) - assert.Equal(t, "hello", lb.String()) + assert.Empty(t, cmp.Diff(11, n)) + assert.Empty(t, cmp.Diff("hello", lb.String())) } func TestLimitedBufferExactFit(t *testing.T) { @@ -132,11 +133,11 @@ func TestLimitedBufferExactFit(t *testing.T) { lb := &limitedBuffer{max: 5} n, err := lb.Write([]byte("hello")) assert.NoError(t, err) - assert.Equal(t, 5, n) - assert.Equal(t, "hello", lb.String()) + assert.Empty(t, cmp.Diff(5, n)) + assert.Empty(t, cmp.Diff("hello", lb.String())) n, err = lb.Write([]byte("more")) assert.NoError(t, err) - assert.Equal(t, 4, n) - assert.Equal(t, "hello", lb.String()) + assert.Empty(t, cmp.Diff(4, n)) + assert.Empty(t, cmp.Diff("hello", lb.String())) } diff --git a/pkg/itr/wasm/transport_test.go b/pkg/itr/wasm/transport_test.go index 0327c1314..11bba66d3 100644 --- a/pkg/itr/wasm/transport_test.go +++ b/pkg/itr/wasm/transport_test.go @@ -5,6 +5,7 @@ import ( "testing" "github.com/ZanzyTHEbar/dragonscale/pkg/itr" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -25,7 +26,7 @@ func TestTransportNonCodeExecForwarded(t *testing.T) { resp, err := transport.Send(t.Context(), req) require.NoError(t, err) assert.True(t, forwarded) - assert.Equal(t, "forwarded", resp.Result) + assert.Empty(t, cmp.Diff("forwarded", resp.Result)) assert.False(t, resp.IsError) } diff --git a/pkg/memory/dag/backfill_test.go b/pkg/memory/dag/backfill_test.go index db0d04243..287eb8eac 100644 --- a/pkg/memory/dag/backfill_test.go +++ b/pkg/memory/dag/backfill_test.go @@ -6,6 +6,7 @@ import ( "time" jsonv2 "github.com/go-json-experiment/json" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -50,15 +51,15 @@ func TestBackfillMissingSessionDAGs_CreatesSnapshotsAndPersistsStatus(t *testing status, err := dag.BackfillMissingSessionDAGs(ctx, d, d.Queries(), agentID, dag.DefaultBackfillOptions()) require.NoError(t, err) require.NotNil(t, status) - assert.Equal(t, 1, status.SnapshotsCreated) - assert.Equal(t, 0, status.Failures) + assert.Empty(t, cmp.Diff(1, status.SnapshotsCreated)) + assert.Empty(t, cmp.Diff(0, status.Failures)) row, err := d.Queries().GetLatestDAGSnapshotBySession(ctx, memsqlc.GetLatestDAGSnapshotBySessionParams{ AgentID: agentID, SessionKey: sessionKey, }) require.NoError(t, err) - assert.Equal(t, int64(12), row.MsgCount) + assert.Empty(t, cmp.Diff(int64(12), row.MsgCount)) kv, err := d.ListKVByPrefix(ctx, agentID, "migration:dag_backfill", 10) require.NoError(t, err) @@ -69,7 +70,7 @@ func TestBackfillMissingSessionDAGs_CreatesSnapshotsAndPersistsStatus(t *testing require.NoError(t, jsonv2.Unmarshal([]byte(raw), &stored)) break } - assert.Equal(t, status.SnapshotsCreated, stored.SnapshotsCreated) + assert.Empty(t, cmp.Diff(status.SnapshotsCreated, stored.SnapshotsCreated)) assert.False(t, stored.CompletedAt.IsZero()) status2, err := dag.BackfillMissingSessionDAGs(ctx, d, d.Queries(), agentID, dag.DefaultBackfillOptions()) diff --git a/pkg/memory/dag/dag_test.go b/pkg/memory/dag/dag_test.go index 73f5b51c4..468de5b12 100644 --- a/pkg/memory/dag/dag_test.go +++ b/pkg/memory/dag/dag_test.go @@ -45,8 +45,8 @@ func TestCompressor_SmallInput(t *testing.T) { require.Len(t, d.Nodes, 1) chunk := d.NodesAtLevel(LevelChunk) require.Len(t, chunk, 1) - assert.Equal(t, 0, chunk[0].StartIdx) - assert.Equal(t, 2, chunk[0].EndIdx) + assert.Empty(t, cmp.Diff(0, chunk[0].StartIdx)) + assert.Empty(t, cmp.Diff(2, chunk[0].EndIdx)) assert.Contains(t, chunk[0].Summary, "user:") assert.Contains(t, chunk[0].Summary, "assistant:") } @@ -63,12 +63,12 @@ func TestCompressor_ChunkSplitting(t *testing.T) { chunks := d.NodesAtLevel(LevelChunk) require.Len(t, chunks, 3) - assert.Equal(t, 0, chunks[0].StartIdx) - assert.Equal(t, 4, chunks[0].EndIdx) - assert.Equal(t, 4, chunks[1].StartIdx) - assert.Equal(t, 8, chunks[1].EndIdx) - assert.Equal(t, 8, chunks[2].StartIdx) - assert.Equal(t, 12, chunks[2].EndIdx) + assert.Empty(t, cmp.Diff(0, chunks[0].StartIdx)) + assert.Empty(t, cmp.Diff(4, chunks[0].EndIdx)) + assert.Empty(t, cmp.Diff(4, chunks[1].StartIdx)) + assert.Empty(t, cmp.Diff(8, chunks[1].EndIdx)) + assert.Empty(t, cmp.Diff(8, chunks[2].StartIdx)) + assert.Empty(t, cmp.Diff(12, chunks[2].EndIdx)) } func TestCompressor_SectionBuilding(t *testing.T) { @@ -89,8 +89,8 @@ func TestCompressor_SectionBuilding(t *testing.T) { assert.Len(t, sections, 3) // First section covers chunks 0-1 (msgs 0-7) - assert.Equal(t, 0, sections[0].StartIdx) - assert.Equal(t, 8, sections[0].EndIdx) + assert.Empty(t, cmp.Diff(0, sections[0].StartIdx)) + assert.Empty(t, cmp.Diff(8, sections[0].EndIdx)) assert.Len(t, sections[0].Children, 2) } @@ -107,12 +107,12 @@ func TestCompressor_SessionSummary(t *testing.T) { sessions := d.NodesAtLevel(LevelSession) require.Len(t, sessions, 1) - assert.Equal(t, 0, sessions[0].StartIdx) - assert.Equal(t, 24, sessions[0].EndIdx) + assert.Empty(t, cmp.Diff(0, sessions[0].StartIdx)) + assert.Empty(t, cmp.Diff(24, sessions[0].EndIdx)) assert.Len(t, sessions[0].Children, 3) require.Len(t, d.Roots, 1) - assert.Equal(t, sessions[0].ID, d.Roots[0]) + assert.Empty(t, cmp.Diff(sessions[0].ID, d.Roots[0])) } func TestExtractSentences(t *testing.T) { @@ -135,7 +135,7 @@ func TestExtractSentences(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { got := extractSentences(tt.text, tt.n) - assert.Equal(t, tt.want, got) + assert.Empty(t, cmp.Diff(tt.want, got)) }) } } @@ -209,8 +209,7 @@ func TestComputeBudget(t *testing.T) { func TestBudget_Remaining(t *testing.T) { t.Parallel() b := Budget{Total: 10000} - assert.Equal(t, 7000, b.Remaining(1000, 500, 500, 500, 500, 0)) - assert.Equal(t, 0, b.Remaining(5000, 3000, 1000, 1000, 1000, 0)) + assert.Empty(t, cmp.Diff(7000, b.Remaining(1000, 500, 500, 500, 500, 0))) } func TestSelectDAGLevel(t *testing.T) { @@ -226,21 +225,21 @@ func TestSelectDAGLevel(t *testing.T) { chunkTokens := d.TotalTokens(LevelChunk) // Large budget -> most detailed (chunk) - assert.Equal(t, LevelChunk, SelectDAGLevel(d, chunkTokens+1000)) + assert.Empty(t, cmp.Diff(LevelChunk, SelectDAGLevel(d, chunkTokens+1000))) // Very small budget -> session level - assert.Equal(t, LevelSession, SelectDAGLevel(d, 10)) + assert.Empty(t, cmp.Diff(LevelSession, SelectDAGLevel(d, 10))) // Nil DAG - assert.Equal(t, LevelRaw, SelectDAGLevel(nil, 1000)) + assert.Empty(t, cmp.Diff(LevelRaw, SelectDAGLevel(nil, 1000))) } func TestTailMessageCount(t *testing.T) { t.Parallel() - assert.Equal(t, 4, TailMessageCount(100)) // Minimum - assert.Equal(t, 20, TailMessageCount(1000)) // 1000/50 - assert.Equal(t, 4, TailMessageCount(0)) // Zero budget - assert.Equal(t, 4, TailMessageCount(-1)) // Negative + assert.Empty(t, cmp.Diff(4, TailMessageCount(100))) // Minimum + assert.Empty(t, cmp.Diff(20, TailMessageCount(1000))) // 1000/50 + assert.Empty(t, cmp.Diff(4, TailMessageCount(0))) // Zero budget + assert.Empty(t, cmp.Diff(4, TailMessageCount(-1))) // Negative } func TestRenderDAGForBudget(t *testing.T) { diff --git a/pkg/memory/delegate/sqlite_audit_test.go b/pkg/memory/delegate/sqlite_audit_test.go index 7fd55a56f..4b4e5bd19 100644 --- a/pkg/memory/delegate/sqlite_audit_test.go +++ b/pkg/memory/delegate/sqlite_audit_test.go @@ -7,6 +7,7 @@ import ( "github.com/ZanzyTHEbar/dragonscale/pkg/ids" "github.com/ZanzyTHEbar/dragonscale/pkg/memory" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -60,7 +61,7 @@ func TestLibSQLDelegate_InsertAuditEntry(t *testing.T) { count, err := d.CountAuditEntries(ctx, tt.entry.AgentID) require.NoError(t, err) - assert.Equal(t, 1, count) + assert.Empty(t, cmp.Diff(1, count)) }) } } @@ -171,9 +172,15 @@ func TestLibSQLDelegate_ListAuditEntriesByAction(t *testing.T) { require.NoError(t, err) assert.Len(t, entries, tt.wantLen) if tt.wantAction != "" { - for _, e := range entries { - assert.Equal(t, tt.wantAction, e.Action) + actions := make([]string, len(entries)) + for i, entry := range entries { + actions[i] = entry.Action } + wantActions := make([]string, len(entries)) + for i := range wantActions { + wantActions[i] = tt.wantAction + } + assert.Empty(t, cmp.Diff(wantActions, actions)) } }) } @@ -234,14 +241,14 @@ func TestLibSQLDelegate_PruneOldAuditEntries(t *testing.T) { count, err := d.CountAuditEntries(ctx, "a1") require.NoError(t, err) - assert.Equal(t, 2, count) + assert.Empty(t, cmp.Diff(2, count)) // Prune entries created before "now + 1 minute" (should remove all) require.NoError(t, d.PruneOldAuditEntries(ctx, "a1", time.Now().Add(time.Minute))) count, err = d.CountAuditEntries(ctx, "a1") require.NoError(t, err) - assert.Equal(t, 0, count) + assert.Empty(t, cmp.Diff(0, count)) } func TestLibSQLDelegate_CountAuditEntriesByAction(t *testing.T) { @@ -281,7 +288,7 @@ func TestLibSQLDelegate_CountAuditEntriesByAction(t *testing.T) { } count, err := d.CountAuditEntriesByAction(ctx, tt.agentID, tt.action) require.NoError(t, err) - assert.Equal(t, tt.want, count) + assert.Empty(t, cmp.Diff(tt.want, count)) }) } } diff --git a/pkg/memory/delegate/sqlite_dag_test.go b/pkg/memory/delegate/sqlite_dag_test.go index 3e527d821..6c8d82881 100644 --- a/pkg/memory/delegate/sqlite_dag_test.go +++ b/pkg/memory/delegate/sqlite_dag_test.go @@ -5,6 +5,8 @@ import ( "github.com/ZanzyTHEbar/dragonscale/pkg/memory/dag" memsqlc "github.com/ZanzyTHEbar/dragonscale/pkg/memory/sqlc" + "github.com/google/go-cmp/cmp" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -38,9 +40,9 @@ func TestLibSQLDelegate_PersistDAG(t *testing.T) { SessionKey: "session1", }) require.NoError(t, err) - require.Equal(t, "agent1", row.AgentID) - require.Equal(t, "session1", row.SessionKey) - require.Equal(t, int64(16), row.MsgCount) + assert.Empty(t, cmp.Diff("agent1", row.AgentID)) + assert.Empty(t, cmp.Diff("session1", row.SessionKey)) + assert.Empty(t, cmp.Diff(int64(16), row.MsgCount)) nodes, err := d.Queries().ListDAGNodesBySnapshotID(ctx, memsqlc.ListDAGNodesBySnapshotIDParams{ SnapshotID: row.ID, diff --git a/pkg/memory/delegate/sqlite_doc_test.go b/pkg/memory/delegate/sqlite_doc_test.go index 989262f4d..c88081f17 100644 --- a/pkg/memory/delegate/sqlite_doc_test.go +++ b/pkg/memory/delegate/sqlite_doc_test.go @@ -6,6 +6,8 @@ import ( "github.com/ZanzyTHEbar/dragonscale/pkg/ids" "github.com/ZanzyTHEbar/dragonscale/pkg/memory" + "github.com/google/go-cmp/cmp" + "github.com/google/go-cmp/cmp/cmpopts" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -79,9 +81,12 @@ func TestLibSQLDelegate_GetDocument(t *testing.T) { return } require.NotNil(t, doc) - assert.Equal(t, tt.wantContent, doc.Content) - assert.Equal(t, tt.agentID, doc.AgentID) - assert.Equal(t, tt.docName, doc.Name) + assert.Empty(t, cmp.Diff(&memory.AgentDocument{ + AgentID: tt.agentID, + Name: tt.docName, + Content: tt.wantContent, + IsActive: true, + }, doc, cmpopts.IgnoreFields(memory.AgentDocument{}, "ID", "Category", "Version", "CreatedAt", "UpdatedAt"))) }) } } @@ -148,7 +153,7 @@ func TestLibSQLDelegate_UpsertDocument(t *testing.T) { got, err := d.GetDocument(ctx, tt.agentID, tt.docName) require.NoError(t, err) require.NotNil(t, got) - assert.Equal(t, tt.wantContent, got.Content) + assert.Empty(t, cmp.Diff(tt.wantContent, got.Content)) }) } } @@ -202,7 +207,7 @@ func TestLibSQLDelegate_DeleteDocument(t *testing.T) { other, err := d.GetDocument(ctx, "a2", "shared") require.NoError(t, err) require.NotNil(t, other, "other agent's document must survive") - assert.Equal(t, "a2-doc", other.Content) + assert.Empty(t, cmp.Diff("a2-doc", other.Content)) } }) } @@ -268,7 +273,7 @@ func TestLibSQLDelegate_ListDocumentsByCategory(t *testing.T) { require.NoError(t, err) assert.Len(t, docs, tt.wantLen) if tt.wantName != "" && len(docs) > 0 { - assert.Equal(t, tt.wantName, docs[0].Name) + assert.Empty(t, cmp.Diff(tt.wantName, docs[0].Name)) } }) } diff --git a/pkg/memory/delegate/sqlite_integration_test.go b/pkg/memory/delegate/sqlite_integration_test.go index 584f7b3df..fea4c4f46 100644 --- a/pkg/memory/delegate/sqlite_integration_test.go +++ b/pkg/memory/delegate/sqlite_integration_test.go @@ -8,10 +8,21 @@ import ( "github.com/ZanzyTHEbar/dragonscale/pkg" "github.com/ZanzyTHEbar/dragonscale/pkg/ids" "github.com/ZanzyTHEbar/dragonscale/pkg/memory" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) +type integrationRecallItemRoleContent struct { + Role string + Content string +} + +type auditEntryActionTarget struct { + Action string + Target string +} + func TestCronKVBackend_Roundtrip(t *testing.T) { t.Parallel() d := newTestDelegate(t) @@ -52,12 +63,18 @@ func TestCronKVBackend_Roundtrip(t *testing.T) { var loaded cronStore require.NoError(t, jsonv2.Unmarshal([]byte(raw), &loaded)) - assert.Equal(t, 1, loaded.Version) - assert.Len(t, loaded.Jobs, 2) - assert.Equal(t, "daily report", loaded.Jobs[0].Name) - assert.True(t, loaded.Jobs[0].Enabled) - assert.Equal(t, "weekly backup", loaded.Jobs[1].Name) - assert.False(t, loaded.Jobs[1].Enabled) + wantStore := cronStore{ + Version: 1, + Jobs: []struct { + ID string `json:"id"` + Name string `json:"name"` + Enabled bool `json:"enabled"` + }{ + {ID: "job-1", Name: "daily report", Enabled: true}, + {ID: "job-2", Name: "weekly backup", Enabled: false}, + }, + } + assert.Empty(t, cmp.Diff(wantStore, loaded)) } func TestCronKVBackend_UpdatePreservesShape(t *testing.T) { @@ -75,7 +92,7 @@ func TestCronKVBackend_UpdatePreservesShape(t *testing.T) { raw, err := d.GetKV(ctx, agentID, kvKey) require.NoError(t, err) - assert.Equal(t, v2, raw) + assert.Empty(t, cmp.Diff(v2, raw)) } func TestCronKVBackend_PrefixScan(t *testing.T) { @@ -111,15 +128,26 @@ func TestEndToEnd_SessionAndAuditFlow(t *testing.T) { // 2. Verify session message count count, err := d.CountSessionMessages(ctx, agentID, sessionKey) require.NoError(t, err) - assert.Equal(t, int64(4), count) + assert.Empty(t, cmp.Diff(int64(4), count)) // 3. Verify message ordering (ASC) msgs, err := d.ListSessionMessages(ctx, agentID, sessionKey, "", 50) require.NoError(t, err) require.Len(t, msgs, 4) - assert.Equal(t, "user", msgs[0].Role) - assert.Equal(t, "What is the weather?", msgs[0].Content) - assert.Equal(t, "assistant", msgs[3].Role) + wantMsgs := []integrationRecallItemRoleContent{ + {Role: "user", Content: "What is the weather?"}, + {Role: "assistant", Content: "Let me check..."}, + {Role: "tool", Content: `{"temp":72,"unit":"F"}`}, + {Role: "assistant", Content: "It's 72F."}, + } + gotMsgs := make([]integrationRecallItemRoleContent, len(msgs)) + for i, msg := range msgs { + gotMsgs[i] = integrationRecallItemRoleContent{ + Role: msg.Role, + Content: msg.Content, + } + } + assert.Empty(t, cmp.Diff(wantMsgs, gotMsgs)) // 4. Insert audit entries for the tool call auditEntry := &memory.AuditEntry{ @@ -138,8 +166,11 @@ func TestEndToEnd_SessionAndAuditFlow(t *testing.T) { auditEntries, err := d.ListAuditEntriesBySession(ctx, agentID, sessionKey, 10) require.NoError(t, err) require.Len(t, auditEntries, 1) - assert.Equal(t, "tool_call", auditEntries[0].Action) - assert.Equal(t, "weather_api", auditEntries[0].Target) + gotAuditEntry := auditEntryActionTarget{ + Action: auditEntries[0].Action, + Target: auditEntries[0].Target, + } + assert.Empty(t, cmp.Diff(auditEntryActionTarget{Action: "tool_call", Target: "weather_api"}, gotAuditEntry)) // 6. Store KV state (e.g. focus checkpoint) require.NoError(t, d.UpsertKV(ctx, agentID, "focus:"+sessionKey, `{"topic":"weather query","checkpoint_index":2}`)) @@ -161,7 +192,7 @@ func TestEndToEnd_SessionAndAuditFlow(t *testing.T) { loadedDoc, err := d.GetDocument(ctx, agentID, "AGENT.md") require.NoError(t, err) require.NotNil(t, loadedDoc) - assert.Equal(t, "# Agent Identity\nI am dragonscale.", loadedDoc.Content) + assert.Empty(t, cmp.Diff("# Agent Identity\nI am dragonscale.", loadedDoc.Content)) // 8. Verify cross-table isolation: different session sees nothing otherMsgs, err := d.ListSessionMessages(ctx, agentID, "sess-other", "", 50) @@ -186,7 +217,7 @@ func TestEndToEnd_WorkingContextAndRecallRoundtrip(t *testing.T) { wc, err := d.GetWorkingContext(ctx, agentID, sessionKey) require.NoError(t, err) require.NotNil(t, wc) - assert.Equal(t, "Initial system prompt state", wc.Content) + assert.Empty(t, cmp.Diff("Initial system prompt state", wc.Content)) // 2. Insert recall items item := &memory.RecallItem{ @@ -206,14 +237,14 @@ func TestEndToEnd_WorkingContextAndRecallRoundtrip(t *testing.T) { // 3. Count recall items recallCount, err := d.CountRecallItems(ctx, agentID, sessionKey) require.NoError(t, err) - assert.Equal(t, 1, recallCount) + assert.Empty(t, cmp.Diff(1, recallCount)) // 4. Update working context require.NoError(t, d.UpsertWorkingContext(ctx, agentID, sessionKey, "Updated with preference awareness")) wc, err = d.GetWorkingContext(ctx, agentID, sessionKey) require.NoError(t, err) - assert.Equal(t, "Updated with preference awareness", wc.Content) + assert.Empty(t, cmp.Diff("Updated with preference awareness", wc.Content)) // 5. Insert a summary summary := &memory.MemorySummary{ @@ -229,5 +260,5 @@ func TestEndToEnd_WorkingContextAndRecallRoundtrip(t *testing.T) { summaries, err := d.ListSummaries(ctx, agentID, sessionKey, 10) require.NoError(t, err) require.Len(t, summaries, 1) - assert.Equal(t, "User discussed preferences. Key info captured.", summaries[0].Content) + assert.Empty(t, cmp.Diff("User discussed preferences. Key info captured.", summaries[0].Content)) } diff --git a/pkg/memory/delegate/sqlite_kv_test.go b/pkg/memory/delegate/sqlite_kv_test.go index 5321639d9..4ac003720 100644 --- a/pkg/memory/delegate/sqlite_kv_test.go +++ b/pkg/memory/delegate/sqlite_kv_test.go @@ -4,6 +4,7 @@ import ( "context" "testing" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -84,7 +85,7 @@ func TestLibSQLDelegate_GetKV(t *testing.T) { return } require.NoError(t, err) - assert.Equal(t, tt.wantValue, got) + assert.Empty(t, cmp.Diff(tt.wantValue, got)) }) } } @@ -155,7 +156,7 @@ func TestLibSQLDelegate_UpsertKV(t *testing.T) { } got, err := d.GetKV(ctx, tt.agentID, tt.key) require.NoError(t, err) - assert.Equal(t, tt.want, got) + assert.Empty(t, cmp.Diff(tt.want, got)) }) } } @@ -214,7 +215,7 @@ func TestLibSQLDelegate_DeleteKV(t *testing.T) { if tt.name == "delete only affects target agent" { other, err := d.GetKV(ctx, "a2", "shared") require.NoError(t, err) - assert.Equal(t, "a2-val", other, "other agent's key must survive") + assert.Empty(t, cmp.Diff("a2-val", other), "other agent's key must survive") } }) } @@ -306,7 +307,7 @@ func TestLibSQLDelegate_ListKVByPrefix(t *testing.T) { return } - assert.Equal(t, tt.want, got) + assert.Empty(t, cmp.Diff(tt.want, got)) }) } } diff --git a/pkg/memory/delegate/sqlite_session_msg_test.go b/pkg/memory/delegate/sqlite_session_msg_test.go index eefcfa8bb..06217e863 100644 --- a/pkg/memory/delegate/sqlite_session_msg_test.go +++ b/pkg/memory/delegate/sqlite_session_msg_test.go @@ -7,6 +7,7 @@ import ( "github.com/ZanzyTHEbar/dragonscale/pkg/ids" "github.com/ZanzyTHEbar/dragonscale/pkg/memory" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -77,11 +78,16 @@ func TestLibSQLDelegate_InsertSessionMessage(t *testing.T) { count, err := d.CountSessionMessages(ctx, tt.agentID, tt.sessionKey) require.NoError(t, err) - assert.Equal(t, int64(1), count) + assert.Empty(t, cmp.Diff(int64(1), count)) }) } } +type recallItemRoleContent struct { + Role string + Content string +} + func TestLibSQLDelegate_ListSessionMessages(t *testing.T) { t.Parallel() tests := []struct { @@ -92,7 +98,8 @@ func TestLibSQLDelegate_ListSessionMessages(t *testing.T) { role string limit int wantLen int - wantRole string + wantRoles []string + wantMsgs []recallItemRoleContent }{ { name: "empty session returns empty", @@ -114,6 +121,11 @@ func TestLibSQLDelegate_ListSessionMessages(t *testing.T) { role: "", limit: 50, wantLen: 3, + wantMsgs: []recallItemRoleContent{ + {Role: "user", Content: "msg1"}, + {Role: "assistant", Content: "msg2"}, + {Role: "tool", Content: "msg3"}, + }, }, { name: "filter by role=user", @@ -127,7 +139,11 @@ func TestLibSQLDelegate_ListSessionMessages(t *testing.T) { role: "user", limit: 50, wantLen: 2, - wantRole: "user", + wantRoles: []string{"user", "user"}, + wantMsgs: []recallItemRoleContent{ + {Role: "user", Content: "u1"}, + {Role: "user", Content: "u2"}, + }, }, { name: "session isolation", @@ -140,6 +156,9 @@ func TestLibSQLDelegate_ListSessionMessages(t *testing.T) { role: "", limit: 50, wantLen: 1, + wantMsgs: []recallItemRoleContent{ + {Role: "user", Content: "msgA"}, + }, }, { name: "agent isolation", @@ -152,6 +171,9 @@ func TestLibSQLDelegate_ListSessionMessages(t *testing.T) { role: "", limit: 50, wantLen: 1, + wantMsgs: []recallItemRoleContent{ + {Role: "user", Content: "from-a1"}, + }, }, { name: "respects limit", @@ -178,6 +200,12 @@ func TestLibSQLDelegate_ListSessionMessages(t *testing.T) { role: "", limit: 50, wantLen: 3, + wantRoles: []string{"user", "assistant", "user"}, + wantMsgs: []recallItemRoleContent{ + {Role: "user", Content: "first"}, + {Role: "assistant", Content: "second"}, + {Role: "user", Content: "third"}, + }, }, } @@ -192,16 +220,22 @@ func TestLibSQLDelegate_ListSessionMessages(t *testing.T) { require.NoError(t, err) assert.Len(t, msgs, tt.wantLen) - if tt.wantRole != "" { - for _, m := range msgs { - assert.Equal(t, tt.wantRole, m.Role) + if tt.wantRoles != nil { + gotRoles := make([]string, len(msgs)) + for i, got := range msgs { + gotRoles[i] = got.Role } + assert.Empty(t, cmp.Diff(tt.wantRoles, gotRoles)) } - - if tt.name == "ordered by created_at ASC" && len(msgs) == 3 { - assert.Equal(t, "first", msgs[0].Content) - assert.Equal(t, "second", msgs[1].Content) - assert.Equal(t, "third", msgs[2].Content) + if tt.wantMsgs != nil { + gotMsgs := make([]recallItemRoleContent, len(msgs)) + for i, got := range msgs { + gotMsgs[i] = recallItemRoleContent{ + Role: got.Role, + Content: got.Content, + } + } + assert.Empty(t, cmp.Diff(tt.wantMsgs, gotMsgs)) } }) } @@ -266,7 +300,7 @@ func TestLibSQLDelegate_CountSessionMessages(t *testing.T) { } count, err := d.CountSessionMessages(ctx, tt.agentID, tt.sessionKey) require.NoError(t, err) - assert.Equal(t, tt.want, count) + assert.Empty(t, cmp.Diff(tt.want, count)) }) } } diff --git a/pkg/memory/integration_test.go b/pkg/memory/integration_test.go index e2bf5c8f6..be4371bed 100644 --- a/pkg/memory/integration_test.go +++ b/pkg/memory/integration_test.go @@ -9,6 +9,7 @@ import ( "github.com/ZanzyTHEbar/dragonscale/pkg/memory" "github.com/ZanzyTHEbar/dragonscale/pkg/memory/delegate" memstore "github.com/ZanzyTHEbar/dragonscale/pkg/memory/store" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -103,8 +104,8 @@ func TestIntegration_BlobPK_RoundTrip(t *testing.T) { fetched, err := store.GetRecall(ctx, item.ID) require.NoError(t, err) require.NotNil(t, fetched) - assert.Equal(t, item.ID, fetched.ID, "BLOB PK round-trip must preserve ID") - assert.Equal(t, "BLOB PK integration test item", fetched.Content) + assert.Empty(t, cmp.Diff(item.ID, fetched.ID), "BLOB PK round-trip must preserve ID") + assert.Empty(t, cmp.Diff("BLOB PK integration test item", fetched.Content)) } func TestIntegration_ArchivalChunking_EndToEnd(t *testing.T) { @@ -129,7 +130,7 @@ func TestIntegration_ArchivalChunking_EndToEnd(t *testing.T) { for _, chunk := range chunks { assert.False(t, chunk.ID.IsZero(), "chunk ID should not be zero") - assert.Equal(t, recallID, chunk.RecallID, "chunk must reference parent recall item") + assert.Empty(t, cmp.Diff(recallID, chunk.RecallID), "chunk must reference parent recall item") } } @@ -226,12 +227,12 @@ func TestIntegration_WorkingContext_Persistence(t *testing.T) { wc, err := del.GetWorkingContext(ctx, testAgent, testSession) require.NoError(t, err) require.NotNil(t, wc) - assert.Equal(t, "initial state", wc.Content) + assert.Empty(t, cmp.Diff("initial state", wc.Content)) require.NoError(t, del.UpsertWorkingContext(ctx, testAgent, testSession, "updated state")) wc, err = del.GetWorkingContext(ctx, testAgent, testSession) require.NoError(t, err) - assert.Equal(t, "updated state", wc.Content) + assert.Empty(t, cmp.Diff("updated state", wc.Content)) } func TestIntegration_Summary_CRUD(t *testing.T) { @@ -252,8 +253,8 @@ func TestIntegration_Summary_CRUD(t *testing.T) { fetched, err := del.ListSummaries(ctx, testAgent, testSession, 1) require.NoError(t, err) require.Len(t, fetched, 1) - assert.Equal(t, summary.ID, fetched[0].ID) - assert.Equal(t, "Summarized conversation about testing", fetched[0].Content) + assert.Empty(t, cmp.Diff(summary.ID, fetched[0].ID)) + assert.Empty(t, cmp.Diff("Summarized conversation about testing", fetched[0].Content)) } func TestIntegration_IDUniqueness_AcrossEntities(t *testing.T) { @@ -300,7 +301,7 @@ func TestIntegration_ContextPressure(t *testing.T) { pressure, err := store.ContextUsage(ctx, testAgent, testSession) require.NoError(t, err) require.NotNil(t, pressure) - assert.Equal(t, 0, pressure.RecallItemCount) + assert.Empty(t, cmp.Diff(0, pressure.RecallItemCount)) for i := 0; i < 3; i++ { item := &memory.RecallItem{ @@ -315,5 +316,5 @@ func TestIntegration_ContextPressure(t *testing.T) { pressure, err = store.ContextUsage(ctx, testAgent, testSession) require.NoError(t, err) - assert.Equal(t, 3, pressure.RecallItemCount) + assert.Empty(t, cmp.Diff(3, pressure.RecallItemCount)) } diff --git a/pkg/memory/kernel_contract_test.go b/pkg/memory/kernel_contract_test.go index e4e2a9757..dbdefe751 100644 --- a/pkg/memory/kernel_contract_test.go +++ b/pkg/memory/kernel_contract_test.go @@ -5,6 +5,7 @@ import ( "time" "github.com/ZanzyTHEbar/dragonscale/pkg/ids" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" ) @@ -17,7 +18,7 @@ func TestActiveContextProjection_TotalTokens(t *testing.T) { {Tokens: 35}, }, } - assert.Equal(t, 235, p.TotalTokens()) + assert.Empty(t, cmp.Diff(235, p.TotalTokens())) } func TestActiveContextProjection_HasLosslessRefs(t *testing.T) { diff --git a/pkg/memory/migrate_sessions_test.go b/pkg/memory/migrate_sessions_test.go index d73a3a1eb..55e92d3e6 100644 --- a/pkg/memory/migrate_sessions_test.go +++ b/pkg/memory/migrate_sessions_test.go @@ -8,6 +8,7 @@ import ( "time" jsonv2 "github.com/go-json-experiment/json" + "github.com/google/go-cmp/cmp" "github.com/ZanzyTHEbar/dragonscale/pkg" "github.com/ZanzyTHEbar/dragonscale/pkg/ids" @@ -114,71 +115,156 @@ func writeSessionFile(t *testing.T, dir, name string, sess SessionFile) { } } -func TestMigrateFileSessions_Basic(t *testing.T) { +func TestMigrateFileSessions(t *testing.T) { t.Parallel() - sessDir := t.TempDir() - del := newMockDelegate() - writeSessionFile(t, sessDir, "sess1.json", SessionFile{ - Key: "session-1", - Messages: []SessionMsg{ - {Role: "user", Content: "Hello"}, - {Role: "assistant", Content: "Hi there!"}, + type tc struct { + name string + setup func(t *testing.T, dir string, del *mockDelegate) + sessionsDir string + want *MigrateSessionsResult + extra func(t *testing.T, dir string, result *MigrateSessionsResult, del *mockDelegate) + } + + tests := []tc{ + { + name: "basic", + setup: func(t *testing.T, dir string, _ *mockDelegate) { + writeSessionFile(t, dir, "sess1.json", SessionFile{ + Key: "session-1", + Messages: []SessionMsg{ + {Role: "user", Content: "Hello"}, + {Role: "assistant", Content: "Hi there!"}, + }, + Created: time.Now().Add(-time.Hour), + Updated: time.Now(), + }) + writeSessionFile(t, dir, "sess2.json", SessionFile{ + Key: "session-2", + Summary: "Talked about Go programming", + Messages: []SessionMsg{ + {Role: "user", Content: "Tell me about Go"}, + }, + }) + }, + want: &MigrateSessionsResult{SessionsFound: 2, SessionsMigrated: 2, ItemsCreated: 3, Errors: 0}, + extra: func(t *testing.T, _ string, _ *MigrateSessionsResult, del *mockDelegate) { + if len(del.recallItems) != 3 { + t.Fatalf("expected 3 recall items, got %d", len(del.recallItems)) + } + first := del.recallItems[0] + if first.Role != "user" || first.Content != "Hello" || first.SessionKey != "session-1" || first.Tags != "migrated" { + t.Errorf("unexpected first recall item: %+v", first) + } + + wc, err := del.GetWorkingContext(t.Context(), pkg.NAME, "session-2") + if err != nil { + t.Fatalf("GetWorkingContext: %v", err) + } + if wc == nil || wc.Content != "Talked about Go programming" { + t.Errorf("expected summary as working context, got %v", wc) + } + }, }, - Created: time.Now().Add(-time.Hour), - Updated: time.Now(), - }) - - writeSessionFile(t, sessDir, "sess2.json", SessionFile{ - Key: "session-2", - Summary: "Talked about Go programming", - Messages: []SessionMsg{ - {Role: "user", Content: "Tell me about Go"}, + { + name: "empty dir", + setup: func(_ *testing.T, _ string, _ *mockDelegate) {}, + want: &MigrateSessionsResult{ + SessionsFound: 0, + SessionsMigrated: 0, + ItemsCreated: 0, + Errors: 0, + }, + }, + { + name: "skip empty messages", + setup: func(t *testing.T, dir string, _ *mockDelegate) { + writeSessionFile(t, dir, "sess.json", SessionFile{ + Key: "s1", + Messages: []SessionMsg{ + {Role: "user", Content: "real content"}, + {Role: "assistant", Content: ""}, + {Role: "user", Content: " "}, + }, + }) + }, + want: &MigrateSessionsResult{SessionsFound: 1, SessionsMigrated: 1, ItemsCreated: 1, Errors: 0}, + }, + { + name: "fallback key from filename", + setup: func(t *testing.T, dir string, _ *mockDelegate) { + writeSessionFile(t, dir, "custom-key.json", SessionFile{ + Key: "", + Messages: []SessionMsg{ + {Role: "user", Content: "test"}, + }, + }) + }, + want: &MigrateSessionsResult{SessionsFound: 1, SessionsMigrated: 1, ItemsCreated: 1, Errors: 0}, + extra: func(t *testing.T, _ string, result *MigrateSessionsResult, del *mockDelegate) { + if got := del.recallItems[0].SessionKey; got != "custom-key" { + t.Errorf("expected 'custom-key' from filename, got %q", got) + } + if result.ItemsCreated != 1 { + t.Fatalf("expected 1 item, got %d", result.ItemsCreated) + } + }, + }, + { + name: "malformed json", + setup: func(t *testing.T, dir string, _ *mockDelegate) { + if err := os.WriteFile(filepath.Join(dir, "bad.json"), []byte("{broken"), 0644); err != nil { + t.Fatalf("write malformed json: %v", err) + } + writeSessionFile(t, dir, "good.json", SessionFile{ + Key: "g1", + Messages: []SessionMsg{ + {Role: "user", Content: "ok"}, + }, + }) + }, + want: &MigrateSessionsResult{SessionsFound: 2, SessionsMigrated: 1, ItemsCreated: 1, Errors: 1}, + }, + { + name: "nonexistent dir", + sessionsDir: "/nonexistent/path", + want: nil, }, - }) - - result, err := MigrateFileSessions(t.Context(), del, pkg.NAME, sessDir) - if err != nil { - t.Fatalf("MigrateFileSessions: %v", err) } - if result.SessionsFound != 2 { - t.Errorf("expected 2 sessions found, got %d", result.SessionsFound) - } - if result.SessionsMigrated != 2 { - t.Errorf("expected 2 sessions migrated, got %d", result.SessionsMigrated) - } - if result.ItemsCreated != 3 { - t.Errorf("expected 3 items, got %d", result.ItemsCreated) - } - if result.Errors != 0 { - t.Errorf("expected 0 errors, got %d", result.Errors) - } + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + del := newMockDelegate() + sessDir := t.TempDir() + if tt.sessionsDir != "" { + sessDir = tt.sessionsDir + } + if tt.setup != nil { + tt.setup(t, sessDir, del) + } - // Check recall items were created with correct data - if len(del.recallItems) != 3 { - t.Fatalf("expected 3 recall items, got %d", len(del.recallItems)) - } - if del.recallItems[0].Role != "user" { - t.Errorf("expected 'user' role, got %q", del.recallItems[0].Role) - } - if del.recallItems[0].Content != "Hello" { - t.Errorf("expected 'Hello', got %q", del.recallItems[0].Content) - } - if del.recallItems[0].SessionKey != "session-1" { - t.Errorf("expected 'session-1', got %q", del.recallItems[0].SessionKey) - } - if del.recallItems[0].Tags != "migrated" { - t.Errorf("expected 'migrated' tag, got %q", del.recallItems[0].Tags) - } - - // Check summary was stored as working context - wc, err := del.GetWorkingContext(t.Context(), pkg.NAME, "session-2") - if err != nil { - t.Fatalf("GetWorkingContext: %v", err) - } - if wc == nil || wc.Content != "Talked about Go programming" { - t.Errorf("expected summary as working context, got %v", wc) + got, err := MigrateFileSessions(t.Context(), del, pkg.NAME, sessDir) + if err != nil { + t.Fatalf("MigrateFileSessions: %v", err) + } + if tt.want == nil { + if got != nil { + t.Errorf("expected nil result, got %#v", got) + } + return + } + if got == nil { + t.Fatal("expected migration result, got nil") + } + if diff := cmp.Diff(*tt.want, *got); diff != "" { + t.Errorf("migration result mismatch (-want +got):\n%s", diff) + } + if tt.extra != nil { + tt.extra(t, sessDir, got, del) + } + }) } } @@ -194,7 +280,6 @@ func TestMigrateFileSessions_Idempotent(t *testing.T) { }, }) - // First run result1, err := MigrateFileSessions(t.Context(), del, pkg.NAME, sessDir) if err != nil { t.Fatalf("first migration: %v", err) @@ -203,7 +288,6 @@ func TestMigrateFileSessions_Idempotent(t *testing.T) { t.Fatalf("expected 1, got %d", result1.SessionsMigrated) } - // Second run should be a no-op (marker file exists) result2, err := MigrateFileSessions(t.Context(), del, pkg.NAME, sessDir) if err != nil { t.Fatalf("second migration: %v", err) @@ -212,114 +296,7 @@ func TestMigrateFileSessions_Idempotent(t *testing.T) { t.Error("expected nil result for already-migrated directory") } - // Still only 1 item if len(del.recallItems) != 1 { t.Errorf("expected 1 recall item (no duplicates), got %d", len(del.recallItems)) } } - -func TestMigrateFileSessions_EmptyDir(t *testing.T) { - t.Parallel() - sessDir := t.TempDir() - del := newMockDelegate() - - result, err := MigrateFileSessions(t.Context(), del, pkg.NAME, sessDir) - if err != nil { - t.Fatalf("MigrateFileSessions: %v", err) - } - if result.SessionsFound != 0 { - t.Errorf("expected 0 sessions, got %d", result.SessionsFound) - } -} - -func TestMigrateFileSessions_NonexistentDir(t *testing.T) { - t.Parallel() - del := newMockDelegate() - - result, err := MigrateFileSessions(t.Context(), del, pkg.NAME, "/nonexistent/path") - if err != nil { - t.Fatalf("MigrateFileSessions: %v", err) - } - if result != nil { - t.Error("expected nil result for nonexistent dir") - } -} - -func TestMigrateFileSessions_SkipsEmptyMessages(t *testing.T) { - t.Parallel() - sessDir := t.TempDir() - del := newMockDelegate() - - writeSessionFile(t, sessDir, "sess.json", SessionFile{ - Key: "s1", - Messages: []SessionMsg{ - {Role: "user", Content: "real content"}, - {Role: "assistant", Content: ""}, - {Role: "user", Content: " "}, - }, - }) - - result, err := MigrateFileSessions(t.Context(), del, pkg.NAME, sessDir) - if err != nil { - t.Fatalf("MigrateFileSessions: %v", err) - } - if result.ItemsCreated != 1 { - t.Errorf("expected 1 item (empty msgs skipped), got %d", result.ItemsCreated) - } -} - -func TestMigrateFileSessions_FallbackKey(t *testing.T) { - t.Parallel() - sessDir := t.TempDir() - del := newMockDelegate() - - // Session with empty key — should use filename - writeSessionFile(t, sessDir, "custom-key.json", SessionFile{ - Key: "", - Messages: []SessionMsg{ - {Role: "user", Content: "test"}, - }, - }) - - result, err := MigrateFileSessions(t.Context(), del, pkg.NAME, sessDir) - if err != nil { - t.Fatalf("MigrateFileSessions: %v", err) - } - if result.ItemsCreated != 1 { - t.Fatalf("expected 1 item, got %d", result.ItemsCreated) - } - if del.recallItems[0].SessionKey != "custom-key" { - t.Errorf("expected 'custom-key' from filename, got %q", del.recallItems[0].SessionKey) - } -} - -func TestMigrateFileSessions_MalformedJSON(t *testing.T) { - t.Parallel() - sessDir := t.TempDir() - del := newMockDelegate() - - // Write a malformed JSON file - os.WriteFile(filepath.Join(sessDir, "bad.json"), []byte("{broken"), 0644) - - // Also a valid one - writeSessionFile(t, sessDir, "good.json", SessionFile{ - Key: "g1", - Messages: []SessionMsg{ - {Role: "user", Content: "ok"}, - }, - }) - - result, err := MigrateFileSessions(t.Context(), del, pkg.NAME, sessDir) - if err != nil { - t.Fatalf("MigrateFileSessions: %v", err) - } - if result.SessionsFound != 2 { - t.Errorf("expected 2 found, got %d", result.SessionsFound) - } - if result.SessionsMigrated != 1 { - t.Errorf("expected 1 migrated, got %d", result.SessionsMigrated) - } - if result.Errors != 1 { - t.Errorf("expected 1 error, got %d", result.Errors) - } -} diff --git a/pkg/memory/observation/observation_test.go b/pkg/memory/observation/observation_test.go index 8a7f173b6..5b6671112 100644 --- a/pkg/memory/observation/observation_test.go +++ b/pkg/memory/observation/observation_test.go @@ -53,17 +53,17 @@ func TestRelativeDate(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { got := relativeDate(tc.ref, now) - assert.Equal(t, tc.want, got) + assert.Empty(t, cmp.Diff(tc.want, got)) }) } } func TestPriorityEmoji(t *testing.T) { t.Parallel() - assert.Equal(t, "🔴", PriorityCritical.Emoji()) - assert.Equal(t, "🟡", PriorityNotable.Emoji()) - assert.Equal(t, "🔵", PriorityInformational.Emoji()) - assert.Equal(t, "🔵", Priority("unknown").Emoji()) + assert.Empty(t, cmp.Diff("🔴", PriorityCritical.Emoji())) + assert.Empty(t, cmp.Diff("🟡", PriorityNotable.Emoji())) + assert.Empty(t, cmp.Diff("🔵", PriorityInformational.Emoji())) + assert.Empty(t, cmp.Diff("🔵", Priority("unknown").Emoji())) } func TestFormatBlock(t *testing.T) { @@ -84,8 +84,8 @@ func TestFormatBlock(t *testing.T) { func TestFormatBlock_Empty(t *testing.T) { t.Parallel() - assert.Equal(t, "", FormatBlock(nil)) - assert.Equal(t, "", FormatBlock([]Observation{})) + assert.Empty(t, cmp.Diff("", FormatBlock(nil))) + assert.Empty(t, cmp.Diff("", FormatBlock([]Observation{}))) } func TestMarshalUnmarshalRoundTrip(t *testing.T) { @@ -185,7 +185,7 @@ func TestObserver_Observe(t *testing.T) { obs, err := o.Observe(t.Context(), msgs, nil) require.NoError(t, err) assert.Len(t, obs, 2) - assert.Equal(t, PriorityCritical, obs[0].Priority) + assert.Empty(t, cmp.Diff(PriorityCritical, obs[0].Priority)) } func TestReflector_ShouldReflect(t *testing.T) { @@ -224,8 +224,8 @@ func TestReflector_Reflect(t *testing.T) { kept, err := r.Reflect(t.Context(), obs) require.NoError(t, err) assert.Len(t, kept, 2) - assert.Equal(t, "Critical fact", kept[0].Content) - assert.Equal(t, "Notable thing", kept[1].Content) + assert.Empty(t, cmp.Diff("Critical fact", kept[0].Content)) + assert.Empty(t, cmp.Diff("Notable thing", kept[1].Content)) } func TestReflector_ReflectFallbackKeepsCritical(t *testing.T) { @@ -245,16 +245,16 @@ func TestReflector_ReflectFallbackKeepsCritical(t *testing.T) { kept, err := r.Reflect(t.Context(), obs) require.NoError(t, err) assert.Len(t, kept, 1) - assert.Equal(t, "Must keep", kept[0].Content) + assert.Empty(t, cmp.Diff("Must keep", kept[0].Content)) } func TestParsePriority(t *testing.T) { t.Parallel() - assert.Equal(t, PriorityCritical, parsePriority("critical")) - assert.Equal(t, PriorityCritical, parsePriority("CRITICAL")) - assert.Equal(t, PriorityNotable, parsePriority("notable")) - assert.Equal(t, PriorityInformational, parsePriority("informational")) - assert.Equal(t, PriorityInformational, parsePriority("unknown")) + assert.Empty(t, cmp.Diff(PriorityCritical, parsePriority("critical"))) + assert.Empty(t, cmp.Diff(PriorityCritical, parsePriority("CRITICAL"))) + assert.Empty(t, cmp.Diff(PriorityNotable, parsePriority("notable"))) + assert.Empty(t, cmp.Diff(PriorityInformational, parsePriority("informational"))) + assert.Empty(t, cmp.Diff(PriorityInformational, parsePriority("unknown"))) } func TestParseKeptIndices(t *testing.T) { diff --git a/pkg/memory/store/chunker_test.go b/pkg/memory/store/chunker_test.go index e8189d51e..a41843afe 100644 --- a/pkg/memory/store/chunker_test.go +++ b/pkg/memory/store/chunker_test.go @@ -4,6 +4,7 @@ import ( "strings" "testing" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -21,7 +22,7 @@ func TestMarkdownChunker_BasicSplit(t *testing.T) { assert.Greater(t, len(chunks), 1, "long content should produce multiple chunks") for i, c := range chunks { - assert.Equal(t, i, c.Index) + assert.Empty(t, cmp.Diff(i, c.Index)) assert.NotEmpty(t, c.Text) } } @@ -33,7 +34,7 @@ func TestMarkdownChunker_SmallContent(t *testing.T) { chunks, err := chunker.Chunk("Short text.") require.NoError(t, err) assert.Len(t, chunks, 1) - assert.Equal(t, "Short text.", chunks[0].Text) + assert.Empty(t, cmp.Diff("Short text.", chunks[0].Text)) } func TestMarkdownChunker_PreservesMarkdownStructure(t *testing.T) { @@ -94,8 +95,8 @@ func TestMarkdownChunker_EmptyContent(t *testing.T) { func TestMarkdownChunker_DefaultConfig(t *testing.T) { t.Parallel() cfg := DefaultMarkdownChunkerConfig() - assert.Equal(t, 1600, cfg.ChunkSize) - assert.Equal(t, 320, cfg.ChunkOverlap) + assert.Empty(t, cmp.Diff(1600, cfg.ChunkSize)) + assert.Empty(t, cmp.Diff(320, cfg.ChunkOverlap)) assert.True(t, cfg.CodeBlocks) assert.True(t, cfg.Headings) } diff --git a/pkg/memory/store/memory_store_test.go b/pkg/memory/store/memory_store_test.go index 173efb018..aa8fb78fa 100644 --- a/pkg/memory/store/memory_store_test.go +++ b/pkg/memory/store/memory_store_test.go @@ -13,6 +13,7 @@ import ( "github.com/ZanzyTHEbar/dragonscale/pkg/memory" memdag "github.com/ZanzyTHEbar/dragonscale/pkg/memory/dag" "github.com/ZanzyTHEbar/dragonscale/pkg/memory/delegate" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -111,7 +112,7 @@ func TestWorkingContext_SetAndGet(t *testing.T) { // Get back content, err = store.GetWorkingContext(ctx, "agent-1", "session-1") require.NoError(t, err) - assert.Equal(t, "You are a helpful assistant.", content) + assert.Empty(t, cmp.Diff("You are a helpful assistant.", content)) // Update err = store.SetWorkingContext(ctx, "agent-1", "session-1", "Updated context.") @@ -119,7 +120,7 @@ func TestWorkingContext_SetAndGet(t *testing.T) { content, err = store.GetWorkingContext(ctx, "agent-1", "session-1") require.NoError(t, err) - assert.Equal(t, "Updated context.", content) + assert.Empty(t, cmp.Diff("Updated context.", content)) } func TestWorkingContext_IsolatedBySessions(t *testing.T) { @@ -134,11 +135,11 @@ func TestWorkingContext_IsolatedBySessions(t *testing.T) { a, err := store.GetWorkingContext(ctx, "agent-1", "session-a") require.NoError(t, err) - assert.Equal(t, "Context A", a) + assert.Empty(t, cmp.Diff("Context A", a)) b, err := store.GetWorkingContext(ctx, "agent-1", "session-b") require.NoError(t, err) - assert.Equal(t, "Context B", b) + assert.Empty(t, cmp.Diff("Context B", b)) } func TestRecall_CRUD(t *testing.T) { @@ -165,8 +166,8 @@ func TestRecall_CRUD(t *testing.T) { got, err := store.GetRecall(ctx, item.ID) require.NoError(t, err) require.NotNil(t, got) - assert.Equal(t, "The user asked about Go generics.", got.Content) - assert.Equal(t, memory.SectorEpisodic, got.Sector) + assert.Empty(t, cmp.Diff("The user asked about Go generics.", got.Content)) + assert.Empty(t, cmp.Diff(memory.SectorEpisodic, got.Sector)) assert.InDelta(t, 0.8, got.Importance, 0.001) // Update @@ -178,7 +179,7 @@ func TestRecall_CRUD(t *testing.T) { updated, err := store.GetRecall(ctx, item.ID) require.NoError(t, err) assert.InDelta(t, 0.95, updated.Importance, 0.001) - assert.Equal(t, "Updated: user asked about Go generics in depth.", updated.Content) + assert.Empty(t, cmp.Diff("Updated: user asked about Go generics in depth.", updated.Content)) // Delete err = store.DeleteRecall(ctx, item.ID) @@ -323,7 +324,7 @@ func TestSearch_ShadowModeUsesBaselineAndTracksParity(t *testing.T) { var metrics retrievalShadowMetrics require.NoError(t, json.Unmarshal([]byte(metricsRaw), &metrics)) - assert.Equal(t, 1, metrics.TotalQueries) + assert.Empty(t, cmp.Diff(1, metrics.TotalQueries)) } func TestSearch_DoesNotPromoteWithoutAugmentedSignals(t *testing.T) { @@ -366,7 +367,7 @@ func TestSearch_DoesNotPromoteWithoutAugmentedSignals(t *testing.T) { require.NoError(t, err) var state retrievalPolicyState require.NoError(t, json.Unmarshal([]byte(stateRaw), &state)) - assert.Equal(t, retrievalModeShadow, state.Mode) + assert.Empty(t, cmp.Diff(retrievalModeShadow, state.Mode)) } func TestSearch_PromoteOnlyOnGateWin(t *testing.T) { @@ -412,7 +413,7 @@ func TestSearch_PromoteOnlyOnGateWin(t *testing.T) { var state retrievalPolicyState require.NoError(t, json.Unmarshal([]byte(stateRaw), &state)) - assert.Equal(t, retrievalModePromoted, state.Mode) + assert.Empty(t, cmp.Diff(retrievalModePromoted, state.Mode)) } func TestSearch_FastRollbackPreservesBaselinePath(t *testing.T) { @@ -488,7 +489,7 @@ func TestSearch_FastRollbackPreservesBaselinePath(t *testing.T) { require.NoError(t, err) var state retrievalPolicyState require.NoError(t, json.Unmarshal([]byte(stateRaw), &state)) - assert.Equal(t, retrievalModeRollback, state.Mode) + assert.Empty(t, cmp.Diff(retrievalModeRollback, state.Mode)) // Subsequent calls should return baseline-only path again. results, err := store.Search(ctx, "grocery", memory.SearchOptions{ @@ -538,7 +539,7 @@ func TestUpdateRetrievalPolicy_PersistFailuresDoNotBlockTransitions(t *testing.T } next := store.updateRetrievalPolicy(ctx, state, gates, metrics, parity, true, baseline) - assert.Equal(t, retrievalModePromoted, next.Mode) + assert.Empty(t, cmp.Diff(retrievalModePromoted, next.Mode)) } func TestSearch_ConcurrentRetrievalPolicyUpdates(t *testing.T) { @@ -596,9 +597,9 @@ func TestContextUsage(t *testing.T) { // Empty system — should be normal pressure pressure, err := store.ContextUsage(ctx, "agent-1", "session-1") require.NoError(t, err) - assert.Equal(t, memory.PressureNormal, pressure.PressureLevel) - assert.Equal(t, 0, pressure.WorkingContextTokens) - assert.Equal(t, 0, pressure.RecallItemCount) + assert.Empty(t, cmp.Diff(memory.PressureNormal, pressure.PressureLevel)) + assert.Empty(t, cmp.Diff(0, pressure.WorkingContextTokens)) + assert.Empty(t, cmp.Diff(0, pressure.RecallItemCount)) // Add working context err = store.SetWorkingContext(ctx, "agent-1", "session-1", strings.Repeat("x", 4000)) @@ -642,7 +643,7 @@ func TestContextUsage_RecallTokenEstimateUsesContent(t *testing.T) { expectedRecallTokens := estimateTokens(contentA) + estimateTokens(contentB) assert.GreaterOrEqual(t, pressure.EstimatedTotalTokens, expectedRecallTokens) - assert.Equal(t, 2, pressure.RecallItemCount) + assert.Empty(t, cmp.Diff(2, pressure.RecallItemCount)) } func TestContextUsage_PressureLevels(t *testing.T) { @@ -782,7 +783,7 @@ func TestRRF_MergesTwoSets(t *testing.T) { merged := ReciprocalRankFusion([][]memory.SearchResult{set1, set2}, []float64{1.0, 1.0}, 60) require.GreaterOrEqual(t, len(merged), 2) // idB appears in both sets, should have highest fused score - assert.Equal(t, idB, merged[0].ID) + assert.Empty(t, cmp.Diff(idB, merged[0].ID)) } func TestRecencyDecay(t *testing.T) { @@ -819,5 +820,5 @@ func TestApplyRecencyDecay_ReordersByAge(t *testing.T) { }) // idNew should now rank higher because idOld got heavily decayed - assert.Equal(t, idNew, results[0].ID) + assert.Empty(t, cmp.Diff(idNew, results[0].ID)) } diff --git a/pkg/memory/store/memory_tool_test.go b/pkg/memory/store/memory_tool_test.go index a266d8ec1..90de0926f 100644 --- a/pkg/memory/store/memory_tool_test.go +++ b/pkg/memory/store/memory_tool_test.go @@ -5,6 +5,7 @@ import ( jsonv2 "github.com/go-json-experiment/json" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -119,8 +120,8 @@ func TestMemoryTool_Status(t *testing.T) { resp := executeAndParse(t, tool, `{"action":"status"}`) assert.True(t, resp.Success) require.NotNil(t, resp.Status) - assert.Equal(t, "normal", resp.Status.PressureLevel) - assert.Equal(t, 0, resp.Status.RecallItemCount) + assert.Empty(t, cmp.Diff("normal", resp.Status.PressureLevel)) + assert.Empty(t, cmp.Diff(0, resp.Status.RecallItemCount)) } func TestMemoryTool_InvalidAction(t *testing.T) { diff --git a/pkg/memory/store/queue_test.go b/pkg/memory/store/queue_test.go index 913205164..76475702d 100644 --- a/pkg/memory/store/queue_test.go +++ b/pkg/memory/store/queue_test.go @@ -6,6 +6,7 @@ import ( "github.com/ZanzyTHEbar/dragonscale/pkg/memory" "github.com/ZanzyTHEbar/dragonscale/pkg/memory/delegate" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -37,7 +38,7 @@ func TestQueueManager_NormalPressure(t *testing.T) { decision, err := qm.Evaluate(ctx, "agent-1", "session-1") require.NoError(t, err) - assert.Equal(t, QueueActionNone, decision.Action) + assert.Empty(t, cmp.Diff(QueueActionNone, decision.Action)) assert.Contains(t, decision.Message, "healthy") } @@ -52,7 +53,7 @@ func TestQueueManager_WarnPressure(t *testing.T) { decision, err := qm.Evaluate(ctx, "agent-1", "s1") require.NoError(t, err) - assert.Equal(t, QueueActionWarn, decision.Action) + assert.Empty(t, cmp.Diff(QueueActionWarn, decision.Action)) assert.Contains(t, decision.Message, "selective") } @@ -67,7 +68,7 @@ func TestQueueManager_OffloadPressure(t *testing.T) { decision, err := qm.Evaluate(ctx, "agent-1", "s1") require.NoError(t, err) - assert.Equal(t, QueueActionOffload, decision.Action) + assert.Empty(t, cmp.Diff(QueueActionOffload, decision.Action)) assert.Contains(t, decision.Message, "offloading") } @@ -82,7 +83,7 @@ func TestQueueManager_FlushPressure(t *testing.T) { decision, err := qm.Evaluate(ctx, "agent-1", "s1") require.NoError(t, err) - assert.Equal(t, QueueActionFlush, decision.Action) + assert.Empty(t, cmp.Diff(QueueActionFlush, decision.Action)) assert.Contains(t, decision.Message, "flush") } @@ -118,7 +119,7 @@ func TestQueueManager_EvictEmpty(t *testing.T) { evicted, summary, err := qm.EvictOldest(ctx, "agent-1", "empty-session") require.NoError(t, err) - assert.Equal(t, 0, evicted) + assert.Empty(t, cmp.Diff(0, evicted)) assert.Empty(t, summary) } @@ -128,5 +129,5 @@ func TestDefaultQueueManagerConfig(t *testing.T) { assert.InDelta(t, 0.70, cfg.WarnThreshold, 0.001) assert.InDelta(t, 0.80, cfg.OffloadThreshold, 0.001) assert.InDelta(t, 0.85, cfg.FlushThreshold, 0.001) - assert.Equal(t, 10, cfg.MaxEvictBatch) + assert.Empty(t, cmp.Diff(10, cfg.MaxEvictBatch)) } diff --git a/pkg/memory/store/scorer_test.go b/pkg/memory/store/scorer_test.go index 4dfd2d93b..e4cc56a3e 100644 --- a/pkg/memory/store/scorer_test.go +++ b/pkg/memory/store/scorer_test.go @@ -4,6 +4,7 @@ import ( "testing" "github.com/ZanzyTHEbar/dragonscale/pkg/memory" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -73,7 +74,7 @@ func TestHeuristicScorer_SectorClassification(t *testing.T) { t.Run(tt.name, func(t *testing.T) { result, err := scorer.Score(ctx, tt.content, "user", "") require.NoError(t, err) - assert.Equal(t, tt.sector, result.Sector) + assert.Empty(t, cmp.Diff(tt.sector, result.Sector)) }) } } @@ -85,7 +86,7 @@ func TestParseScoringResponse_ValidJSON(t *testing.T) { require.NoError(t, err) assert.InDelta(t, 0.85, result.Importance, 0.001) assert.InDelta(t, 0.6, result.Salience, 0.001) - assert.Equal(t, memory.SectorSemantic, result.Sector) + assert.Empty(t, cmp.Diff(memory.SectorSemantic, result.Sector)) } func TestParseScoringResponse_WithCodeFences(t *testing.T) { @@ -94,7 +95,7 @@ func TestParseScoringResponse_WithCodeFences(t *testing.T) { result, err := parseScoringResponse(input) require.NoError(t, err) assert.InDelta(t, 0.9, result.Importance, 0.001) - assert.Equal(t, memory.SectorProcedural, result.Sector) + assert.Empty(t, cmp.Diff(memory.SectorProcedural, result.Sector)) } func TestParseScoringResponse_ClampsValues(t *testing.T) { @@ -111,7 +112,7 @@ func TestParseScoringResponse_UnknownSector(t *testing.T) { input := `{"importance": 0.5, "salience": 0.5, "sector": "unknown_sector"}` result, err := parseScoringResponse(input) require.NoError(t, err) - assert.Equal(t, memory.SectorEpisodic, result.Sector, "unknown sector should default to episodic") + assert.Empty(t, cmp.Diff(memory.SectorEpisodic, result.Sector), "unknown sector should default to episodic") } func TestNormalizeSector(t *testing.T) { @@ -129,7 +130,7 @@ func TestNormalizeSector(t *testing.T) { } for _, tt := range tests { t.Run(tt.input, func(t *testing.T) { - assert.Equal(t, tt.expected, normalizeSector(tt.input)) + assert.Empty(t, cmp.Diff(tt.expected, normalizeSector(tt.input))) }) } } diff --git a/pkg/rlm/engine_test.go b/pkg/rlm/engine_test.go index 559b74c4a..31a94e633 100644 --- a/pkg/rlm/engine_test.go +++ b/pkg/rlm/engine_test.go @@ -7,6 +7,7 @@ import ( "testing" "github.com/ZanzyTHEbar/dragonscale/pkg/rlm" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -110,8 +111,8 @@ func TestFanOut_AllPartitions_Processed(t *testing.T) { assert.Len(t, results, 4) for i, r := range results { - assert.Equal(t, i, r.PartitionIdx) - assert.Equal(t, "ans-"+partitions[i], r.Answer) + assert.Empty(t, cmp.Diff(i, r.PartitionIdx)) + assert.Empty(t, cmp.Diff("ans-"+partitions[i], r.Answer)) } } @@ -126,7 +127,7 @@ func TestFanOut_UnboundedConcurrency(t *testing.T) { return rlm.PartitionResult{PartitionIdx: idx, Answer: part, Tokens: 2} }) assert.Len(t, results, 20) - assert.Equal(t, uint32(40), rlm.TotalTokens(results)) + assert.Empty(t, cmp.Diff(uint32(40), rlm.TotalTokens(results))) } func TestFanOut_Empty(t *testing.T) { @@ -147,7 +148,7 @@ func TestMergeResults_Deduplication(t *testing.T) { {Answer: ""}, // empty — should be skipped } merged := rlm.MergeResults(results) - assert.Equal(t, "alpha\nbeta", merged) + assert.Empty(t, cmp.Diff("alpha\nbeta", merged)) } func TestMergeResults_WithErrors(t *testing.T) { @@ -157,7 +158,7 @@ func TestMergeResults_WithErrors(t *testing.T) { {Err: fmt.Errorf("failed"), Answer: "should be skipped"}, } merged := rlm.MergeResults(results) - assert.Equal(t, "good", merged) + assert.Empty(t, cmp.Diff("good", merged)) } func TestStrategyPlanner_SmallContext_OpFinal(t *testing.T) { @@ -167,7 +168,7 @@ func TestStrategyPlanner_SmallContext_OpFinal(t *testing.T) { planner := rlm.NewStrategyPlanner(cfg) op := planner.PlanNext(500, "any query", 0) - assert.Equal(t, rlm.OpFinal, op.Type) + assert.Empty(t, cmp.Diff(rlm.OpFinal, op.Type)) } func TestStrategyPlanner_MaxDepth_OpFinal(t *testing.T) { @@ -177,7 +178,7 @@ func TestStrategyPlanner_MaxDepth_OpFinal(t *testing.T) { planner := rlm.NewStrategyPlanner(cfg) op := planner.PlanNext(100000, "any query", 3) // depth == MaxDepth - assert.Equal(t, rlm.OpFinal, op.Type) + assert.Empty(t, cmp.Diff(rlm.OpFinal, op.Type)) } func TestStrategyPlanner_KeywordQuery_OpGrep(t *testing.T) { @@ -186,7 +187,7 @@ func TestStrategyPlanner_KeywordQuery_OpGrep(t *testing.T) { // "find" prefix should trigger grep. op := planner.PlanNext(100000, "find myFunction in code", 0) - assert.Equal(t, rlm.OpGrep, op.Type) + assert.Empty(t, cmp.Diff(rlm.OpGrep, op.Type)) assert.NotEmpty(t, op.GrepQuery) } @@ -195,6 +196,6 @@ func TestStrategyPlanner_LargeContext_OpPartition(t *testing.T) { planner := rlm.NewStrategyPlanner(rlm.DefaultStrategyConfig()) op := planner.PlanNext(100000, "summarise everything", 0) - assert.Equal(t, rlm.OpPartition, op.Type) + assert.Empty(t, cmp.Diff(rlm.OpPartition, op.Type)) assert.Greater(t, op.PartitionK, 0) } diff --git a/pkg/rlm/fanout_test.go b/pkg/rlm/fanout_test.go index 64aec1921..94cd83db7 100644 --- a/pkg/rlm/fanout_test.go +++ b/pkg/rlm/fanout_test.go @@ -7,6 +7,7 @@ import ( "testing" "time" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -32,9 +33,9 @@ func TestFanOutUnbounded(t *testing.T) { require.Len(t, results, 3) for i, r := range results { - assert.Equal(t, i, r.PartitionIdx) + assert.Empty(t, cmp.Diff(i, r.PartitionIdx)) assert.Contains(t, r.Answer, fmt.Sprintf("part-%d", i)) - assert.Equal(t, uint32(10), r.Tokens) + assert.Empty(t, cmp.Diff(uint32(10), r.Tokens)) } } @@ -86,8 +87,8 @@ func TestFanOutPreservesOrder(t *testing.T) { require.Len(t, results, 4) for i, r := range results { - assert.Equal(t, i, r.PartitionIdx) - assert.Equal(t, partitions[i], r.Answer) + assert.Empty(t, cmp.Diff(i, r.PartitionIdx)) + assert.Empty(t, cmp.Diff(partitions[i], r.Answer)) } } @@ -102,7 +103,7 @@ func TestMergeResultsDeduplication(t *testing.T) { } merged := MergeResults(results) - assert.Equal(t, "answer one\nanswer two", merged) + assert.Empty(t, cmp.Diff("answer one\nanswer two", merged)) } func TestMergeResultsAllErrors(t *testing.T) { @@ -130,10 +131,10 @@ func TestTotalTokens(t *testing.T) { {Tokens: 250}, {Tokens: 50}, } - assert.Equal(t, uint32(400), TotalTokens(results)) + assert.Empty(t, cmp.Diff(uint32(400), TotalTokens(results))) } func TestTotalTokensEmpty(t *testing.T) { t.Parallel() - assert.Equal(t, uint32(0), TotalTokens(nil)) + assert.Empty(t, cmp.Diff(uint32(0), TotalTokens(nil))) } diff --git a/pkg/rlm/rope_test.go b/pkg/rlm/rope_test.go index 51557d3cd..9b33bf1de 100644 --- a/pkg/rlm/rope_test.go +++ b/pkg/rlm/rope_test.go @@ -5,6 +5,7 @@ import ( "testing" "github.com/ZanzyTHEbar/dragonscale/pkg/rlm" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -12,24 +13,24 @@ import ( func TestRope_EmptyRope(t *testing.T) { t.Parallel() r := rlm.NewRope("") - assert.Equal(t, 0, r.Len()) - assert.Equal(t, "", r.String()) + assert.Empty(t, cmp.Diff(0, r.Len())) + assert.Empty(t, cmp.Diff("", r.String())) } func TestRope_BasicAppendAndString(t *testing.T) { t.Parallel() r := rlm.NewRope("hello") r.Append(" world") - assert.Equal(t, 11, r.Len()) - assert.Equal(t, "hello world", r.String()) + assert.Empty(t, cmp.Diff(11, r.Len())) + assert.Empty(t, cmp.Diff("hello world", r.String())) } func TestRope_LargeContent(t *testing.T) { t.Parallel() content := strings.Repeat("abcdefghij", 1000) // 10000 bytes r := rlm.NewRope(content) - assert.Equal(t, 10000, r.Len()) - assert.Equal(t, content, r.String()) + assert.Empty(t, cmp.Diff(10000, r.Len())) + assert.Empty(t, cmp.Diff(content, r.String())) } func TestRope_Slice_ValidRange(t *testing.T) { @@ -37,7 +38,7 @@ func TestRope_Slice_ValidRange(t *testing.T) { r := rlm.NewRope("hello world") s, err := r.Slice(6, 11) require.NoError(t, err) - assert.Equal(t, "world", s) + assert.Empty(t, cmp.Diff("world", s)) } func TestRope_Slice_ZeroLength(t *testing.T) { @@ -45,7 +46,7 @@ func TestRope_Slice_ZeroLength(t *testing.T) { r := rlm.NewRope("hello") s, err := r.Slice(2, 2) require.NoError(t, err) - assert.Equal(t, "", s) + assert.Empty(t, cmp.Diff("", s)) } func TestRope_Slice_OutOfRange(t *testing.T) { @@ -61,14 +62,14 @@ func TestRope_Slice_AcrossAppendBoundary(t *testing.T) { r.Append(" world") s, err := r.Slice(3, 8) require.NoError(t, err) - assert.Equal(t, "lo wo", s) + assert.Empty(t, cmp.Diff("lo wo", s)) } func TestRope_Lines(t *testing.T) { t.Parallel() r := rlm.NewRope("line1\nline2\nline3") lines := r.Lines() - assert.Equal(t, []string{"line1", "line2", "line3"}, lines) + assert.Empty(t, cmp.Diff([]string{"line1", "line2", "line3"}, lines)) } func TestRope_GrepLines_CaseSensitive(t *testing.T) { @@ -76,8 +77,8 @@ func TestRope_GrepLines_CaseSensitive(t *testing.T) { r := rlm.NewRope("apple\nBanana\napricot\ncherry") matches := r.GrepLines("ap", 0, false) require.Len(t, matches, 2) - assert.Equal(t, 1, matches[0].LineNum) - assert.Equal(t, 3, matches[1].LineNum) + assert.Empty(t, cmp.Diff(1, matches[0].LineNum)) + assert.Empty(t, cmp.Diff(3, matches[1].LineNum)) } func TestRope_GrepLines_CaseInsensitive(t *testing.T) { @@ -85,7 +86,7 @@ func TestRope_GrepLines_CaseInsensitive(t *testing.T) { r := rlm.NewRope("Apple\nbanana\nAPRICOT") matches := r.GrepLines("apple", 0, true) require.Len(t, matches, 1) - assert.Equal(t, "Apple", matches[0].Line) + assert.Empty(t, cmp.Diff("Apple", matches[0].Line)) } func TestRope_GrepLines_MaxMatches(t *testing.T) { @@ -107,7 +108,7 @@ func TestRope_Partition_Even(t *testing.T) { r := rlm.NewRope("12345678") parts := r.Partition(4) assert.Len(t, parts, 4) - assert.Equal(t, "12345678", strings.Join(parts, "")) + assert.Empty(t, cmp.Diff("12345678", strings.Join(parts, ""))) } func TestRope_Partition_MoreThanContent(t *testing.T) { @@ -117,7 +118,7 @@ func TestRope_Partition_MoreThanContent(t *testing.T) { assert.Len(t, parts, 10) // All content should appear in first non-empty partition. combined := strings.Join(parts, "") - assert.Equal(t, "hi", combined) + assert.Empty(t, cmp.Diff("hi", combined)) } func TestRope_Partition_Empty(t *testing.T) { @@ -126,7 +127,7 @@ func TestRope_Partition_Empty(t *testing.T) { parts := r.Partition(4) assert.Len(t, parts, 4) for _, p := range parts { - assert.Equal(t, "", p) + assert.Empty(t, cmp.Diff("", p)) } } @@ -136,6 +137,6 @@ func TestRope_RuneLen(t *testing.T) { ) r := rlm.NewRope("héllo") // 'é' is 2 bytes - assert.Equal(t, 5, r.RuneLen()) - assert.Equal(t, 6, r.Len()) // bytes + assert.Empty(t, cmp.Diff(5, r.RuneLen())) + assert.Empty(t, cmp.Diff(6, r.Len())) // bytes } diff --git a/pkg/rlm/strategy_test.go b/pkg/rlm/strategy_test.go index 4a6b31b29..8663db0ab 100644 --- a/pkg/rlm/strategy_test.go +++ b/pkg/rlm/strategy_test.go @@ -3,6 +3,7 @@ package rlm import ( "testing" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" ) @@ -15,7 +16,7 @@ func TestStrategyPlanNextFinalAtMaxDepth(t *testing.T) { }) op := sp.PlanNext(100000, "any query", 3) - assert.Equal(t, OpFinal, op.Type) + assert.Empty(t, cmp.Diff(OpFinal, op.Type)) } func TestStrategyPlanNextFinalSmallContext(t *testing.T) { @@ -23,7 +24,7 @@ func TestStrategyPlanNextFinalSmallContext(t *testing.T) { sp := NewStrategyPlanner(DefaultStrategyConfig()) op := sp.PlanNext(1000, "any query", 0) - assert.Equal(t, OpFinal, op.Type) + assert.Empty(t, cmp.Diff(OpFinal, op.Type)) } func TestStrategyPlanNextGrepForKeywordQuery(t *testing.T) { @@ -43,7 +44,7 @@ func TestStrategyPlanNextGrepForKeywordQuery(t *testing.T) { for _, tt := range tests { t.Run(tt.query, func(t *testing.T) { op := sp.PlanNext(1_000_000, tt.query, 0) - assert.Equal(t, OpGrep, op.Type) + assert.Empty(t, cmp.Diff(OpGrep, op.Type)) assert.NotEmpty(t, op.GrepQuery) }) } @@ -53,31 +54,31 @@ func TestStrategyPlanNextPartitionDefault(t *testing.T) { t.Parallel() sp := NewStrategyPlanner(DefaultStrategyConfig()) op := sp.PlanNext(100_000, "summarize this document", 0) - assert.Equal(t, OpPartition, op.Type) - assert.Equal(t, 4, op.PartitionK) + assert.Empty(t, cmp.Diff(OpPartition, op.Type)) + assert.Empty(t, cmp.Diff(4, op.PartitionK)) } func TestStrategyPlanNextPartitionLargeContext(t *testing.T) { t.Parallel() sp := NewStrategyPlanner(DefaultStrategyConfig()) op := sp.PlanNext(5_000_000, "summarize this corpus", 0) - assert.Equal(t, OpPartition, op.Type) - assert.Equal(t, 8, op.PartitionK, "large contexts should use more partitions") + assert.Empty(t, cmp.Diff(OpPartition, op.Type)) + assert.Empty(t, cmp.Diff(8, op.PartitionK), "large contexts should use more partitions") } func TestExtractKeywordQuoted(t *testing.T) { t.Parallel() - assert.Equal(t, "handleRequest", extractKeyword(`find "handleRequest" in the codebase`)) + assert.Empty(t, cmp.Diff("handleRequest", extractKeyword(`find "handleRequest" in the codebase`))) } func TestExtractKeywordNoQuotes(t *testing.T) { t.Parallel() - assert.Equal(t, "find", extractKeyword("find the main function")) + assert.Empty(t, cmp.Diff("find", extractKeyword("find the main function"))) } func TestExtractKeywordEmpty(t *testing.T) { t.Parallel() - assert.Equal(t, "", extractKeyword("")) + assert.Empty(t, cmp.Diff("", extractKeyword(""))) } func TestLooksLikeKeywordQuery(t *testing.T) { diff --git a/pkg/runtime/runtime_test.go b/pkg/runtime/runtime_test.go index ea73c792a..26f75c115 100644 --- a/pkg/runtime/runtime_test.go +++ b/pkg/runtime/runtime_test.go @@ -9,6 +9,7 @@ import ( "github.com/ZanzyTHEbar/dragonscale/pkg" "github.com/ZanzyTHEbar/dragonscale/pkg/bus" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -40,11 +41,10 @@ func TestResolveBaseConfigPath_PrefersXDGOverLegacy(t *testing.T) { require.NoError(t, os.WriteFile(legacyPath, []byte(`{}`), 0o644)) got := ResolveBaseConfigPath() - assert.Equal(t, xdgPath, got) + assert.Empty(t, cmp.Diff(xdgPath, got)) } func TestResolveBaseConfigPath_FallsBackToLegacyWhenXDGMissing(t *testing.T) { - t.Parallel() home := t.TempDir() xdg := t.TempDir() t.Setenv("HOME", home) @@ -55,7 +55,7 @@ func TestResolveBaseConfigPath_FallsBackToLegacyWhenXDGMissing(t *testing.T) { require.NoError(t, os.WriteFile(legacyPath, []byte(`{}`), 0o644)) got := ResolveBaseConfigPath() - assert.Equal(t, legacyPath, got) + assert.Empty(t, cmp.Diff(legacyPath, got)) } func TestLoadResolvedConfig_AppliesOverlayAndKeepsBaseValues(t *testing.T) { @@ -79,7 +79,7 @@ func TestLoadResolvedConfig_AppliesOverlayAndKeepsBaseValues(t *testing.T) { OverlayConfigPath: overlayPath, }) require.NoError(t, err) - assert.Equal(t, "base-key", cfg.Providers.OpenAI.APIKey) + assert.Empty(t, cmp.Diff("base-key", cfg.Providers.OpenAI.APIKey)) assert.True(t, cfg.Agents.Defaults.RestrictToSandbox) } @@ -94,7 +94,7 @@ func TestEnsureMinProviderTimeout_SetsFloor(t *testing.T) { MinProviderTimeout: 180 * time.Second, }) require.NoError(t, err) - assert.Equal(t, 180, cfg.Providers.OpenAI.Timeout) + assert.Empty(t, cmp.Diff(180, cfg.Providers.OpenAI.Timeout)) } func TestStartOutbound_DropAndConsumeDoNotBlockPublishers(t *testing.T) { @@ -156,7 +156,7 @@ func TestStartOutbound_CallbackReceivesMessages(t *testing.T) { select { case got := <-received: - assert.Equal(t, "hello", got.Content) + assert.Empty(t, cmp.Diff("hello", got.Content)) case <-time.After(2 * time.Second): t.Fatal("did not receive callback outbound message") } diff --git a/pkg/security/jsonextract_test.go b/pkg/security/jsonextract_test.go index 352e6c635..df265ebdf 100644 --- a/pkg/security/jsonextract_test.go +++ b/pkg/security/jsonextract_test.go @@ -4,6 +4,7 @@ import ( "strings" "testing" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -25,7 +26,7 @@ func TestExtractJSON_RawJSON(t *testing.T) { var result map[string]interface{} err := ExtractJSON(tc.input, &result, nil) require.NoError(t, err) - assert.Equal(t, tc.wantVal, result[tc.wantKey]) + assert.Empty(t, cmp.Diff(tc.wantVal, result[tc.wantKey])) }) } } @@ -54,7 +55,7 @@ func TestExtractJSON_CodeFence(t *testing.T) { var result map[string]interface{} err := ExtractJSON(tc.input, &result, nil) require.NoError(t, err) - assert.Equal(t, float64(42), result["score"]) + assert.Empty(t, cmp.Diff(float64(42), result["score"])) }) } } @@ -68,8 +69,8 @@ func TestExtractJSON_EmbeddedInProse(t *testing.T) { } err := ExtractJSON(input, &result, nil) require.NoError(t, err) - assert.Equal(t, 0.9, result.Importance) - assert.Equal(t, "semantic", result.Sector) + assert.Empty(t, cmp.Diff(0.9, result.Importance)) + assert.Empty(t, cmp.Diff("semantic", result.Sector)) } func TestExtractJSON_NestedBracesInStrings(t *testing.T) { @@ -78,8 +79,8 @@ func TestExtractJSON_NestedBracesInStrings(t *testing.T) { var result map[string]interface{} err := ExtractJSON(input, &result, nil) require.NoError(t, err) - assert.Equal(t, "function() { return {}; }", result["content"]) - assert.Equal(t, float64(1), result["count"]) + assert.Empty(t, cmp.Diff("function() { return {}; }", result["content"])) + assert.Empty(t, cmp.Diff(float64(1), result["count"])) } func TestExtractJSON_InjectionAttempts(t *testing.T) { @@ -148,7 +149,7 @@ func TestExtractJSON_MultipleFences_TakesFirst(t *testing.T) { var result map[string]interface{} err := ExtractJSON(input, &result, nil) require.NoError(t, err) - assert.Equal(t, true, result["first"]) + assert.Empty(t, cmp.Diff(true, result["first"])) _, hasSecond := result["second"] assert.False(t, hasSecond) } @@ -168,9 +169,9 @@ func TestSanitizeToolArgs_ValidInput(t *testing.T) { result, err := SanitizeToolArgs(args, schema) require.NoError(t, err) - assert.Equal(t, "/tmp/test.txt", result["path"]) - assert.Equal(t, "hello world", result["content"]) - assert.Equal(t, "overwrite", result["mode"]) + assert.Empty(t, cmp.Diff("/tmp/test.txt", result["path"])) + assert.Empty(t, cmp.Diff("hello world", result["content"])) + assert.Empty(t, cmp.Diff("overwrite", result["mode"])) } func TestSanitizeToolArgs_MissingRequired(t *testing.T) { @@ -224,7 +225,7 @@ func TestSanitizeToolArgs_TypeCoercion(t *testing.T) { assert.Error(t, err) } else { require.NoError(t, err) - assert.Equal(t, tc.expected, result["arg"]) + assert.Empty(t, cmp.Diff(tc.expected, result["arg"])) } }) } @@ -234,7 +235,7 @@ func TestExtractFirstBraced_EscapedQuotes(t *testing.T) { t.Parallel() input := `{"msg": "say \"hello\" world"}` result := extractFirstBraced(input) - assert.Equal(t, input, result) + assert.Empty(t, cmp.Diff(input, result)) } func TestExtractFirstBraced_UnbalancedBraces(t *testing.T) { diff --git a/pkg/security/redact_test.go b/pkg/security/redact_test.go index be00fb01e..d55b42743 100644 --- a/pkg/security/redact_test.go +++ b/pkg/security/redact_test.go @@ -3,6 +3,7 @@ package security import ( "testing" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" ) @@ -92,7 +93,7 @@ func TestRedactor_SafeText(t *testing.T) { t.Parallel() r := NewRedactor() safe := "This is a normal log message about processing 42 items." - assert.Equal(t, safe, r.Redact(safe)) + assert.Empty(t, cmp.Diff(safe, r.Redact(safe))) assert.False(t, r.ContainsSensitive(safe)) } @@ -109,10 +110,10 @@ func TestRedactor_RedactMap(t *testing.T) { } out := r.RedactMap(m) assert.Contains(t, out["command"].(string), "[REDACTED:") - assert.Equal(t, "normal output", out["output"]) + assert.Empty(t, cmp.Diff("normal output", out["output"])) nested := out["nested"].(map[string]interface{}) assert.Contains(t, nested["secret"].(string), "[REDACTED:") - assert.Equal(t, 42, out["count"]) + assert.Empty(t, cmp.Diff(42, out["count"])) } func TestMaskKey(t *testing.T) { @@ -128,7 +129,7 @@ func TestMaskKey(t *testing.T) { } for _, tc := range tests { t.Run(tc.input, func(t *testing.T) { - assert.Equal(t, tc.want, MaskKey(tc.input)) + assert.Empty(t, cmp.Diff(tc.want, MaskKey(tc.input))) }) } } diff --git a/pkg/security/securebus/audit_test.go b/pkg/security/securebus/audit_test.go index 9e5668016..883442cc9 100644 --- a/pkg/security/securebus/audit_test.go +++ b/pkg/security/securebus/audit_test.go @@ -6,6 +6,7 @@ import ( "testing" "time" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -24,12 +25,12 @@ func TestAuditLogAppendAndRetrieve(t *testing.T) { } require.NoError(t, al.Append(event)) - assert.Equal(t, 1, al.Len()) + assert.Empty(t, cmp.Diff(1, al.Len())) events := al.Events() require.Len(t, events, 1) - assert.Equal(t, "req-1", events[0].RequestID) - assert.Equal(t, "read_file", events[0].ToolName) + assert.Empty(t, cmp.Diff("req-1", events[0].RequestID)) + assert.Empty(t, cmp.Diff("read_file", events[0].ToolName)) } func TestAuditLogConcurrentAppend(t *testing.T) { @@ -50,7 +51,7 @@ func TestAuditLogConcurrentAppend(t *testing.T) { } wg.Wait() - assert.Equal(t, n, al.Len()) + assert.Empty(t, cmp.Diff(n, al.Len())) } func TestAuditLogFilterBySession(t *testing.T) { @@ -81,8 +82,8 @@ func TestAuditLogLeakEvents(t *testing.T) { leaks := al.LeakEvents() assert.Len(t, leaks, 2) - assert.Equal(t, "r2", leaks[0].RequestID) - assert.Equal(t, "r3", leaks[1].RequestID) + assert.Empty(t, cmp.Diff("r2", leaks[0].RequestID)) + assert.Empty(t, cmp.Diff("r3", leaks[1].RequestID)) } type mockSink struct { @@ -109,7 +110,7 @@ func TestAuditLogSinkIntegration(t *testing.T) { _ = al.Append(AuditEvent{RequestID: "r1"}) _ = al.Append(AuditEvent{RequestID: "r2"}) - assert.Equal(t, 2, al.Len()) + assert.Empty(t, cmp.Diff(2, al.Len())) assert.Len(t, sink.events, 2) } @@ -123,7 +124,7 @@ func TestAuditLogSinkError(t *testing.T) { err := al.Append(AuditEvent{RequestID: "r2"}) assert.Error(t, err) - assert.Equal(t, 2, al.Len(), "in-memory log should always append") + assert.Empty(t, cmp.Diff(2, al.Len()), "in-memory log should always append") } func TestAuditLogEventsImmutable(t *testing.T) { @@ -135,5 +136,5 @@ func TestAuditLogEventsImmutable(t *testing.T) { events[0].RequestID = "mutated" original := al.Events() - assert.Equal(t, "r1", original[0].RequestID, "original should be unaffected") + assert.Empty(t, cmp.Diff("r1", original[0].RequestID), "original should be unaffected") } diff --git a/pkg/security/securebus/bus_test.go b/pkg/security/securebus/bus_test.go index cff31e2b6..65c162a0c 100644 --- a/pkg/security/securebus/bus_test.go +++ b/pkg/security/securebus/bus_test.go @@ -10,6 +10,7 @@ import ( "github.com/ZanzyTHEbar/dragonscale/pkg/security" "github.com/ZanzyTHEbar/dragonscale/pkg/security/securebus" "github.com/ZanzyTHEbar/dragonscale/pkg/tools" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -86,8 +87,8 @@ func TestBus_SuccessfulToolExec(t *testing.T) { resp := bus.Execute(t.Context(), req) assert.False(t, resp.IsError) - assert.Equal(t, "hello world", resp.Result) - assert.Equal(t, 1, bus.AuditLog().Len()) + assert.Empty(t, cmp.Diff("hello world", resp.Result)) + assert.Empty(t, cmp.Diff(1, bus.AuditLog().Len())) } func TestBus_UnknownTool(t *testing.T) { @@ -111,7 +112,7 @@ func TestBus_ToolReturnsError(t *testing.T) { resp := bus.Execute(t.Context(), req) assert.True(t, resp.IsError) - assert.Equal(t, 1, bus.AuditLog().Len()) + assert.Empty(t, cmp.Diff(1, bus.AuditLog().Len())) events := bus.AuditLog().Events() assert.True(t, events[0].IsError) } @@ -187,7 +188,7 @@ func TestBus_SecretInjection_ArgVariant(t *testing.T) { resp := bus.Execute(t.Context(), req) assert.False(t, resp.IsError) - assert.Equal(t, "supersecret", resp.Result, "injected secret should appear in tool output") + assert.Empty(t, cmp.Diff("supersecret", resp.Result), "injected secret should appear in tool output") events := bus.AuditLog().Events() require.Len(t, events, 1) @@ -222,7 +223,7 @@ func TestBus_AuditLog_FilterBySession(t *testing.T) { bus.Execute(t.Context(), req) } - assert.Equal(t, 3, bus.AuditLog().Len()) + assert.Empty(t, cmp.Diff(3, bus.AuditLog().Len())) assert.Len(t, bus.AuditLog().FilterBySession("session-A"), 2) assert.Len(t, bus.AuditLog().FilterBySession("session-B"), 1) } @@ -237,7 +238,7 @@ func TestBus_Transport_Send(t *testing.T) { resp, err := bus.Transport().Send(t.Context(), req) require.NoError(t, err) - assert.Equal(t, "pong", resp.Result) + assert.Empty(t, cmp.Diff("pong", resp.Result)) } func TestBus_InvalidArgsJSON(t *testing.T) { @@ -261,7 +262,7 @@ func TestBus_RLMFinalCommand(t *testing.T) { resp := bus.Execute(t.Context(), req) assert.False(t, resp.IsError) - assert.Equal(t, "the answer", resp.Result) + assert.Empty(t, cmp.Diff("the answer", resp.Result)) } func TestBus_CloseIdempotent(t *testing.T) { diff --git a/pkg/security/securebus/socket_transport_test.go b/pkg/security/securebus/socket_transport_test.go index ca654ca14..9ca0c0c18 100644 --- a/pkg/security/securebus/socket_transport_test.go +++ b/pkg/security/securebus/socket_transport_test.go @@ -9,6 +9,7 @@ import ( "time" "github.com/ZanzyTHEbar/dragonscale/pkg/itr" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -50,7 +51,7 @@ func TestSocketTransportRoundTrip(t *testing.T) { resp, err := client.Send(t.Context(), req) require.NoError(t, err) - assert.Equal(t, "req-001", resp.ID) + assert.Empty(t, cmp.Diff("req-001", resp.ID)) assert.Contains(t, resp.Result, "req-001") server.Close() @@ -86,7 +87,7 @@ func TestSocketTransportMultipleRequests(t *testing.T) { ) resp, err := client.Send(t.Context(), req) require.NoError(t, err) - assert.Equal(t, req.ID, resp.ID) + assert.Empty(t, cmp.Diff(req.ID, resp.ID)) } } diff --git a/pkg/security/urlguard_test.go b/pkg/security/urlguard_test.go index 1815faa34..f3979f5f8 100644 --- a/pkg/security/urlguard_test.go +++ b/pkg/security/urlguard_test.go @@ -4,6 +4,7 @@ import ( "net" "testing" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" ) @@ -121,7 +122,7 @@ func TestIsBlockedIP(t *testing.T) { if ip == nil { t.Fatalf("invalid IP: %s", tc.ip) } - assert.Equal(t, tc.blocked, isBlockedIP(ip)) + assert.Empty(t, cmp.Diff(tc.blocked, isBlockedIP(ip))) }) } } diff --git a/pkg/security/vault_test.go b/pkg/security/vault_test.go index e70fa6cc2..2e0e72658 100644 --- a/pkg/security/vault_test.go +++ b/pkg/security/vault_test.go @@ -3,6 +3,7 @@ package security import ( "testing" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -30,7 +31,7 @@ func TestVault_RoundTrip(t *testing.T) { dec, err := v.DecryptString(enc) require.NoError(t, err) - assert.Equal(t, secret, dec) + assert.Empty(t, cmp.Diff(secret, dec)) }) } } @@ -92,7 +93,7 @@ func TestVault_EmptyInput(t *testing.T) { dec, err := v.DecryptString(enc) require.NoError(t, err) - assert.Equal(t, "", dec) + assert.Empty(t, cmp.Diff("", dec)) } func TestVault_BinaryData(t *testing.T) { @@ -106,5 +107,5 @@ func TestVault_BinaryData(t *testing.T) { dec, err := v.Decrypt(enc) require.NoError(t, err) - assert.Equal(t, binary, dec) + assert.Empty(t, cmp.Diff(binary, dec)) } diff --git a/pkg/security/zkp_test.go b/pkg/security/zkp_test.go index 320f67876..e60c06f62 100644 --- a/pkg/security/zkp_test.go +++ b/pkg/security/zkp_test.go @@ -4,6 +4,7 @@ import ( "testing" "time" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -95,7 +96,7 @@ func TestZKPSessionManagerIssueAndValidate(t *testing.T) { st, err := sm.VerifyAndIssue(commit.RX, commit.RY, challenge, response) require.NoError(t, err) assert.True(t, st.IsValid()) - assert.Equal(t, 1, sm.ActiveSessions()) + assert.Empty(t, cmp.Diff(1, sm.ActiveSessions())) assert.True(t, sm.ValidateToken(st.Token)) } @@ -111,7 +112,7 @@ func TestZKPSessionManagerRejectsInvalidProof(t *testing.T) { _, err := sm.VerifyAndIssue(make([]byte, 32), make([]byte, 32), make([]byte, 32), make([]byte, 32)) assert.Error(t, err) - assert.Equal(t, 0, sm.ActiveSessions()) + assert.Empty(t, cmp.Diff(0, sm.ActiveSessions())) } func TestZKPSessionManagerExpiry(t *testing.T) { @@ -152,7 +153,7 @@ func TestZKPSessionManagerRevoke(t *testing.T) { sm.RevokeToken(st.Token) assert.False(t, sm.ValidateToken(st.Token)) - assert.Equal(t, 0, sm.ActiveSessions()) + assert.Empty(t, cmp.Diff(0, sm.ActiveSessions())) } func TestHandshakePayloadBinaryRoundTrip(t *testing.T) { @@ -166,11 +167,11 @@ func TestHandshakePayloadBinaryRoundTrip(t *testing.T) { } data := hp.MarshalBinary() - assert.Equal(t, 129, len(data)) + assert.Empty(t, cmp.Diff(129, len(data))) decoded, err := UnmarshalBinaryHandshake(data) require.NoError(t, err) - assert.Equal(t, hp, decoded) + assert.Empty(t, cmp.Diff(hp, decoded)) } func TestHandshakeResultBinaryRoundTrip(t *testing.T) { @@ -181,11 +182,11 @@ func TestHandshakeResultBinaryRoundTrip(t *testing.T) { } data := hr.MarshalBinary() - assert.Equal(t, 40, len(data)) + assert.Empty(t, cmp.Diff(40, len(data))) decoded, err := UnmarshalBinaryResult(data) require.NoError(t, err) - assert.Equal(t, hr, decoded) + assert.Empty(t, cmp.Diff(hr, decoded)) } func TestZKPSessionManagerCleanup(t *testing.T) { @@ -203,10 +204,10 @@ func TestZKPSessionManagerCleanup(t *testing.T) { response, _ := ProverRespond(commit, challenge, x) _, _ = sm.VerifyAndIssue(commit.RX, commit.RY, challenge, response) } - assert.Equal(t, 5, sm.ActiveSessions()) + assert.Empty(t, cmp.Diff(5, sm.ActiveSessions())) time.Sleep(5 * time.Millisecond) cleaned := sm.Cleanup() - assert.Equal(t, 5, cleaned) - assert.Equal(t, 0, sm.ActiveSessions()) + assert.Empty(t, cmp.Diff(5, cleaned)) + assert.Empty(t, cmp.Diff(0, sm.ActiveSessions())) } diff --git a/pkg/session/manager_test.go b/pkg/session/manager_test.go index 58c999b2d..45cf092f3 100644 --- a/pkg/session/manager_test.go +++ b/pkg/session/manager_test.go @@ -324,7 +324,7 @@ func TestSessionManager_ProjectionPointerPersistedAndRestored(t *testing.T) { assert.NotZero(t, ptr.LastMessageID) assert.False(t, ptr.FirstCreatedAt.IsZero()) assert.False(t, ptr.LastCreatedAt.IsZero()) - assert.Equal(t, 3, ptr.Count) + assert.Empty(t, cmp.Diff(3, ptr.Count)) // New manager restores; pointer is re-persisted (same values) sm2 := NewSessionManager("", WithSessionDelegate(del, "test-agent")) @@ -355,7 +355,7 @@ func TestSessionManager_ProjectionPointerUpdatedOnAppend(t *testing.T) { require.NotEmpty(t, raw) var ptr ProjectionPointer require.NoError(t, jsonv2.Unmarshal([]byte(raw), &ptr)) - assert.Equal(t, i+1, ptr.Count) + assert.Empty(t, cmp.Diff(i+1, ptr.Count)) } } @@ -380,8 +380,8 @@ func TestSessionManager_IntegrityMismatchRestoreStillSucceeds(t *testing.T) { sm2 := NewSessionManager("", WithSessionDelegate(del, "test-agent")) history := sm2.GetHistory(sessionKey) require.Len(t, history, 2) - assert.Equal(t, "a", history[0].Content) - assert.Equal(t, "b", history[1].Content) + assert.Empty(t, cmp.Diff("a", history[0].Content)) + assert.Empty(t, cmp.Diff("b", history[1].Content)) // Pointer should now reflect restored state raw, err := del.GetKV(t.Context(), "test-agent", projectionPointerKey(sessionKey)) @@ -389,7 +389,7 @@ func TestSessionManager_IntegrityMismatchRestoreStillSucceeds(t *testing.T) { require.NotEmpty(t, raw) var ptr ProjectionPointer require.NoError(t, jsonv2.Unmarshal([]byte(raw), &ptr)) - assert.Equal(t, 2, ptr.Count) + assert.Empty(t, cmp.Diff(2, ptr.Count)) } func TestSessionManager_ProjectionBackfillStatusPersistedOnBootstrap(t *testing.T) { @@ -412,7 +412,7 @@ func TestSessionManager_ProjectionBackfillStatusPersistedOnBootstrap(t *testing. var status ProjectionBackfillStatus require.NoError(t, jsonv2.Unmarshal([]byte(raw), &status)) - assert.Equal(t, 1, status.Version) + assert.Empty(t, cmp.Diff(1, status.Version)) assert.GreaterOrEqual(t, status.SessionsScanned, 2) assert.GreaterOrEqual(t, status.PointersUpdated, 0) assert.WithinDuration(t, time.Now().UTC(), status.CompletedAt, 5*time.Second) diff --git a/pkg/skills/graph_test.go b/pkg/skills/graph_test.go index 40ea6c1fc..4a4a69736 100644 --- a/pkg/skills/graph_test.go +++ b/pkg/skills/graph_test.go @@ -5,6 +5,7 @@ import ( "path/filepath" "testing" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -61,7 +62,7 @@ func TestParseWikilinks(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { got := ParseWikilinks(tc.content) - assert.Equal(t, tc.want, got) + assert.Empty(t, cmp.Diff(tc.want, got)) }) } } @@ -83,7 +84,7 @@ func TestMergeUnique(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { got := mergeUnique(tc.a, tc.b) - assert.Equal(t, tc.want, got) + assert.Empty(t, cmp.Diff(tc.want, got)) }) } } @@ -140,8 +141,8 @@ No wikilinks here. rm := g.GetNode("risk-management") require.NotNil(t, rm) - assert.Equal(t, []string{"trading", "risk"}, rm.Tags) - assert.Equal(t, "finance", rm.Domain) + assert.Empty(t, cmp.Diff([]string{"trading", "risk"}, rm.Tags)) + assert.Empty(t, cmp.Diff("finance", rm.Domain)) assert.Contains(t, rm.Links, "position-sizing") assert.Contains(t, rm.Links, "technical-analysis") @@ -277,7 +278,7 @@ Content. results = g.SearchSkills("software engineering") require.True(t, len(results) >= 1) - assert.Equal(t, "code-review", results[0].Name) + assert.Empty(t, cmp.Diff("code-review", results[0].Name)) results = g.SearchSkills("") assert.Empty(t, results) @@ -307,7 +308,7 @@ Content. mocs := g.ListMOCs() assert.Len(t, mocs, 1) - assert.Equal(t, "trading-moc", mocs[0].Name) + assert.Empty(t, cmp.Diff("trading-moc", mocs[0].Name)) assert.True(t, mocs[0].IsMOC) } @@ -333,7 +334,7 @@ Content. idx := g.GetIndex() require.NotNil(t, idx) - assert.Equal(t, "index", idx.Name) + assert.Empty(t, cmp.Diff("index", idx.Name)) assert.True(t, idx.IsIndex) assert.Contains(t, idx.Links, "trading-moc") assert.Contains(t, idx.Links, "engineering-moc") @@ -369,9 +370,9 @@ Content with [[some-link]]. require.Len(t, skills, 1) s := skills[0] - assert.Equal(t, "json-skill", s.Name) - assert.Equal(t, []string{"alpha", "beta"}, s.Tags) - assert.Equal(t, "testing", s.Domain) + assert.Empty(t, cmp.Diff("json-skill", s.Name)) + assert.Empty(t, cmp.Diff([]string{"alpha", "beta"}, s.Tags)) + assert.Empty(t, cmp.Diff("testing", s.Domain)) } func nodeNames(nodes []*SkillNode) []string { diff --git a/pkg/skills/templates_test.go b/pkg/skills/templates_test.go index 1b3aacec5..1df147810 100644 --- a/pkg/skills/templates_test.go +++ b/pkg/skills/templates_test.go @@ -5,6 +5,7 @@ import ( "path/filepath" "testing" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -54,7 +55,7 @@ func TestInstallTemplate_BuildsValidGraph(t *testing.T) { rm := g.GetNode("risk-management") require.NotNil(t, rm) - assert.Equal(t, "finance", rm.Domain) + assert.Empty(t, cmp.Diff("finance", rm.Domain)) assert.Contains(t, rm.Tags, "trading") } diff --git a/pkg/sync/identity_test.go b/pkg/sync/identity_test.go index 1a1dabf77..8d5caa3ca 100644 --- a/pkg/sync/identity_test.go +++ b/pkg/sync/identity_test.go @@ -9,6 +9,7 @@ import ( "time" "github.com/ZanzyTHEbar/dragonscale/pkg/memory" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -101,7 +102,7 @@ func TestSyncAll_InsertsNewFiles(t *testing.T) { for _, name := range IdentityFiles { doc := store.getDoc("agent-1", name) require.NotNil(t, doc, "expected document for %s", name) - assert.Equal(t, syncCategory, doc.Category) + assert.Empty(t, cmp.Diff(syncCategory, doc.Category)) assert.NotEmpty(t, doc.Content) hash := store.getHash("agent-1", name) @@ -125,7 +126,7 @@ func TestSyncAll_SkipsUnchangedFiles(t *testing.T) { require.NoError(t, s.SyncAll(t.Context())) secondDoc := store.getDoc("agent-1", "AGENT.md") - assert.Equal(t, firstID, secondDoc.ID, "unchanged file should not be re-upserted") + assert.Empty(t, cmp.Diff(firstID, secondDoc.ID), "unchanged file should not be re-upserted") } func TestSyncAll_UpsertsModifiedFiles(t *testing.T) { @@ -235,7 +236,7 @@ func TestCheckAndSync_SkipsUntouchedFiles(t *testing.T) { require.NoError(t, s.CheckAndSync(t.Context())) hash2 := store.getHash("agent-1", "AGENT.md") - assert.Equal(t, hash1, hash2, "untouched file should not trigger re-sync") + assert.Empty(t, cmp.Diff(hash1, hash2), "untouched file should not trigger re-sync") } func TestWatch_DetectsFileChange(t *testing.T) { @@ -285,7 +286,7 @@ func TestContentHash_Deterministic(t *testing.T) { data := []byte("hello world") h1 := contentHash(data) h2 := contentHash(data) - assert.Equal(t, h1, h2) + assert.Empty(t, cmp.Diff(h1, h2)) assert.Len(t, h1, 64) } @@ -313,7 +314,7 @@ func TestIsIdentityFile(t *testing.T) { } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - assert.Equal(t, tt.want, isIdentityFile(tt.name)) + assert.Empty(t, cmp.Diff(tt.want, isIdentityFile(tt.name))) }) } } @@ -322,7 +323,7 @@ func TestNew_SetsFields(t *testing.T) { t.Parallel() store := newMockStore() s := New("/tmp/identity", "test-agent", store) - assert.Equal(t, "/tmp/identity", s.identityDir) - assert.Equal(t, "test-agent", s.agentID) + assert.Empty(t, cmp.Diff("/tmp/identity", s.identityDir)) + assert.Empty(t, cmp.Diff("test-agent", s.agentID)) assert.NotNil(t, s.store) } diff --git a/pkg/tools/agentic_map_test.go b/pkg/tools/agentic_map_test.go index 750b1ad8f..e62bd4c4e 100644 --- a/pkg/tools/agentic_map_test.go +++ b/pkg/tools/agentic_map_test.go @@ -6,6 +6,7 @@ import ( "testing" jsonv2 "github.com/go-json-experiment/json" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -41,9 +42,9 @@ func TestAgenticMapTool_Execute_Success(t *testing.T) { } `json:"summary"` } require.NoError(t, jsonv2.Unmarshal([]byte(result.ForLLM), &payload)) - assert.Equal(t, 2, payload.Count) - assert.Equal(t, 2, payload.Summary.SuccessCount) - assert.Equal(t, 0, payload.Summary.FailureCount) + assert.Empty(t, cmp.Diff(2, payload.Count)) + assert.Empty(t, cmp.Diff(2, payload.Summary.SuccessCount)) + assert.Empty(t, cmp.Diff(0, payload.Summary.FailureCount)) } func TestAgenticMapTool_Execute_RetriesFailedItems(t *testing.T) { @@ -78,7 +79,7 @@ func TestAgenticMapTool_Execute_RetriesFailedItems(t *testing.T) { require.NoError(t, jsonv2.Unmarshal([]byte(result.ForLLM), &payload)) require.Len(t, payload.Results, 1) assert.True(t, payload.Results[0].Success) - assert.Equal(t, 2, payload.Results[0].Attempts) + assert.Empty(t, cmp.Diff(2, payload.Results[0].Attempts)) } func TestAgenticMapTool_Execute_RequiresPlaceholders(t *testing.T) { diff --git a/pkg/tools/focus_test.go b/pkg/tools/focus_test.go index 55a9c5fc8..b995418f0 100644 --- a/pkg/tools/focus_test.go +++ b/pkg/tools/focus_test.go @@ -5,6 +5,7 @@ import ( "testing" jsonv2 "github.com/go-json-experiment/json" + "github.com/google/go-cmp/cmp" "github.com/ZanzyTHEbar/dragonscale/pkg/messages" "github.com/ZanzyTHEbar/dragonscale/pkg/session" @@ -78,8 +79,8 @@ func TestStartFocus(t *testing.T) { var state FocusState require.NoError(t, jsonv2.Unmarshal([]byte(raw), &state)) - assert.Equal(t, "investigate auth bug", state.Topic) - assert.Equal(t, 2, state.CheckpointIndex) + assert.Empty(t, cmp.Diff("investigate auth bug", state.Topic)) + assert.Empty(t, cmp.Diff(2, state.CheckpointIndex)) } func TestStartFocus_MissingTopic(t *testing.T) { @@ -115,7 +116,7 @@ func TestCompleteFocus(t *testing.T) { sm.AddMessage(sk, "assistant", "all tests pass") historyBefore := sm.GetHistory(sk) - require.Equal(t, 8, len(historyBefore)) + assert.Empty(t, cmp.Diff(8, len(historyBefore))) completeTool := NewCompleteFocusTool(delegate, sm, func() string { return sk }) result := completeTool.Execute(ctx, map[string]interface{}{ @@ -130,8 +131,8 @@ func TestCompleteFocus(t *testing.T) { assert.Less(t, len(historyAfter), len(historyBefore)) // Pre-checkpoint messages should be preserved - assert.Equal(t, "hello", historyAfter[0].Content) - assert.Equal(t, "hi", historyAfter[1].Content) + assert.Empty(t, cmp.Diff("hello", historyAfter[0].Content)) + assert.Empty(t, cmp.Diff("hi", historyAfter[1].Content)) // Knowledge should be persisted knowledgeRaw, _ := delegate.GetKV(ctx, focusAgentID, knowledgeKVPrefix+sk) @@ -140,7 +141,7 @@ func TestCompleteFocus(t *testing.T) { var kb KnowledgeBlock require.NoError(t, jsonv2.Unmarshal([]byte(knowledgeRaw), &kb)) require.Len(t, kb.Entries, 1) - assert.Equal(t, "debug auth", kb.Entries[0].Topic) + assert.Empty(t, cmp.Diff("debug auth", kb.Entries[0].Topic)) assert.Contains(t, kb.Entries[0].Summary, "token validation") // Focus state should be cleaned up @@ -192,8 +193,8 @@ func TestCompleteFocus_MultipleKnowledgeEntries(t *testing.T) { var kb KnowledgeBlock require.NoError(t, jsonv2.Unmarshal([]byte(knowledgeRaw), &kb)) require.Len(t, kb.Entries, 2) - assert.Equal(t, "topic A", kb.Entries[0].Topic) - assert.Equal(t, "topic B", kb.Entries[1].Topic) + assert.Empty(t, cmp.Diff("topic A", kb.Entries[0].Topic)) + assert.Empty(t, cmp.Diff("topic B", kb.Entries[1].Topic)) } func TestPruneHistory(t *testing.T) { @@ -255,7 +256,7 @@ func TestPruneHistory(t *testing.T) { // Pre-checkpoint messages should always be preserved for i := 0; i < tt.checkpointIdx && i < len(result); i++ { - assert.Equal(t, tt.history[i].Content, result[i].Content) + assert.Empty(t, cmp.Diff(tt.history[i].Content, result[i].Content)) } }) } diff --git a/pkg/tools/llm_map_test.go b/pkg/tools/llm_map_test.go index 6759b5dfe..752e7229f 100644 --- a/pkg/tools/llm_map_test.go +++ b/pkg/tools/llm_map_test.go @@ -6,6 +6,7 @@ import ( "testing" jsonv2 "github.com/go-json-experiment/json" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -85,9 +86,9 @@ func TestLLMMapTool_Execute_Success(t *testing.T) { Results []map[string]interface{} `json:"results"` } require.NoError(t, jsonv2.Unmarshal([]byte(result.ForLLM), &payload)) - assert.Equal(t, 2, payload.Count) + assert.Empty(t, cmp.Diff(2, payload.Count)) require.Len(t, payload.Results, 2) - assert.Equal(t, "alpha", payload.Results[0]["label"]) + assert.Empty(t, cmp.Diff("alpha", payload.Results[0]["label"])) } func TestLLMMapTool_Execute_SchemaValidationFailure(t *testing.T) { diff --git a/pkg/tools/map_runtime_integration_test.go b/pkg/tools/map_runtime_integration_test.go index 38d9288a2..f48e74300 100644 --- a/pkg/tools/map_runtime_integration_test.go +++ b/pkg/tools/map_runtime_integration_test.go @@ -8,6 +8,7 @@ import ( fantasy "charm.land/fantasy" jsonv2 "github.com/go-json-experiment/json" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -110,13 +111,13 @@ func TestLLMMap_WorkerJSONL_StatusAndRead(t *testing.T) { })) runID, _ := enqueue["run_id"].(string) require.NotEmpty(t, runID) - assert.Equal(t, "worker", enqueue["execution_mode"]) + assert.Empty(t, cmp.Diff("worker", enqueue["execution_mode"])) status := decodeResultMap(t, statusTool.Execute(t.Context(), map[string]interface{}{ "run_id": runID, "process_steps": float64(20), })) - assert.Equal(t, mapRunStatusSucceeded, status["status"]) + assert.Empty(t, cmp.Diff(mapRunStatusSucceeded, status["status"])) assert.EqualValues(t, 2, status["succeeded_items"]) readJSON := decodeResultMap(t, readTool.Execute(t.Context(), map[string]interface{}{ @@ -158,8 +159,8 @@ func TestLLMMap_IdempotencyReuse(t *testing.T) { first := decodeResultMap(t, mapTool.Execute(t.Context(), args)) second := decodeResultMap(t, mapTool.Execute(t.Context(), args)) - assert.Equal(t, first["run_id"], second["run_id"]) - assert.Equal(t, true, second["idempotent_reuse"]) + assert.Empty(t, cmp.Diff(first["run_id"], second["run_id"])) + assert.Empty(t, cmp.Diff(true, second["idempotent_reuse"])) } func TestLLMMap_IdempotencyReuse_Concurrent(t *testing.T) { @@ -248,7 +249,7 @@ func TestLLMMap_InlineRetriesThenSucceeds(t *testing.T) { "execution_mode": "inline", "max_retries": float64(2), })) - assert.Equal(t, mapRunStatusSucceeded, result["status"]) + assert.Empty(t, cmp.Diff(mapRunStatusSucceeded, result["status"])) runID, _ := result["run_id"].(string) require.NotEmpty(t, runID) @@ -261,7 +262,7 @@ func TestLLMMap_InlineRetriesThenSucceeds(t *testing.T) { require.Len(t, itemsAny, 1) item0 := itemsAny[0].(map[string]interface{}) assert.EqualValues(t, 2, item0["attempts"]) - assert.Equal(t, mapItemStatusSucceeded, item0["status"]) + assert.Empty(t, cmp.Diff(mapItemStatusSucceeded, item0["status"])) } func TestLLMMap_InlineExhaustedRetriesFailsRun(t *testing.T) { @@ -282,7 +283,7 @@ func TestLLMMap_InlineExhaustedRetriesFailsRun(t *testing.T) { "execution_mode": "inline", "max_retries": float64(1), })) - assert.Equal(t, mapRunStatusFailed, payload["status"]) + assert.Empty(t, cmp.Diff(mapRunStatusFailed, payload["status"])) summary, ok := payload["summary"].(map[string]interface{}) require.True(t, ok) assert.EqualValues(t, 1, summary["failure_count"]) @@ -360,7 +361,7 @@ func TestAgenticMap_WorkerLifecycle(t *testing.T) { "run_id": runID, "process_steps": float64(20), })) - assert.Equal(t, mapRunStatusSucceeded, status["status"]) + assert.Empty(t, cmp.Diff(mapRunStatusSucceeded, status["status"])) read := decodeResultMap(t, readTool.Execute(t.Context(), map[string]interface{}{ "run_id": runID, @@ -370,7 +371,7 @@ func TestAgenticMap_WorkerLifecycle(t *testing.T) { require.True(t, ok) require.Len(t, itemsAny, 1) item0 := itemsAny[0].(map[string]interface{}) - assert.Equal(t, mapItemStatusSucceeded, item0["status"]) + assert.Empty(t, cmp.Diff(mapItemStatusSucceeded, item0["status"])) } func TestLLMMap_InvalidJSONLIngestFails(t *testing.T) { diff --git a/pkg/tools/obligation_test.go b/pkg/tools/obligation_test.go index a31c673e6..f727cbf1c 100644 --- a/pkg/tools/obligation_test.go +++ b/pkg/tools/obligation_test.go @@ -5,6 +5,7 @@ import ( "time" jsonv2 "github.com/go-json-experiment/json" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -32,7 +33,7 @@ func TestObligationTool_CreateAndList(t *testing.T) { var rec ObligationRecord require.NoError(t, jsonv2.Unmarshal([]byte(create.ForLLM), &rec)) require.NotEmpty(t, rec.ID) - assert.Equal(t, ObligationStateScheduled, rec.State) + assert.Empty(t, cmp.Diff(ObligationStateScheduled, rec.State)) list := tool.Execute(ctx, map[string]interface{}{"action": "list"}) require.NotNil(t, list) @@ -110,7 +111,7 @@ func TestObligationTool_StateMachineAndEvidence(t *testing.T) { var verified ObligationRecord require.NoError(t, jsonv2.Unmarshal([]byte(toVerified.ForLLM), &verified)) - assert.Equal(t, ObligationStateVerified, verified.State) + assert.Empty(t, cmp.Diff(ObligationStateVerified, verified.State)) assert.NotZero(t, verified.VerifiedAt) require.Len(t, verified.Evidence, 1) } @@ -133,15 +134,15 @@ func TestObligationTool_CollectDueObligations_TransitionsScheduledToDue(t *testi var created ObligationRecord require.NoError(t, jsonv2.Unmarshal([]byte(create.ForLLM), &created)) - require.Equal(t, ObligationStateScheduled, created.State) + assert.Empty(t, cmp.Diff(ObligationStateScheduled, created.State)) due, err := tool.CollectDueObligations(ctx, time.Now().UTC(), "heartbeat") require.NoError(t, err) require.Len(t, due, 1) - assert.Equal(t, created.ID, due[0].ID) - assert.Equal(t, ObligationStateDue, due[0].State) + assert.Empty(t, cmp.Diff(created.ID, due[0].ID)) + assert.Empty(t, cmp.Diff(ObligationStateDue, due[0].State)) require.NotEmpty(t, due[0].Evidence) - assert.Equal(t, "heartbeat", due[0].Evidence[0].Source) + assert.Empty(t, cmp.Diff("heartbeat", due[0].Evidence[0].Source)) get := tool.Execute(ctx, map[string]interface{}{ "action": "get", @@ -151,7 +152,7 @@ func TestObligationTool_CollectDueObligations_TransitionsScheduledToDue(t *testi var persisted ObligationRecord require.NoError(t, jsonv2.Unmarshal([]byte(get.ForLLM), &persisted)) - assert.Equal(t, ObligationStateDue, persisted.State) + assert.Empty(t, cmp.Diff(ObligationStateDue, persisted.State)) require.Len(t, persisted.Evidence, 1) } @@ -182,5 +183,5 @@ func TestObligationTool_CollectDueObligations_DoesNotDuplicateDueEvidence(t *tes require.NoError(t, err) require.Len(t, second, 1) require.Len(t, second[0].Evidence, 1, "due transition evidence should be appended only once") - assert.Equal(t, created.ID, second[0].ID) + assert.Empty(t, cmp.Diff(created.ID, second[0].ID)) } diff --git a/pkg/tools/retrieval_test.go b/pkg/tools/retrieval_test.go index 615b01212..a7361d2f3 100644 --- a/pkg/tools/retrieval_test.go +++ b/pkg/tools/retrieval_test.go @@ -3,6 +3,7 @@ package tools import ( "testing" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -10,7 +11,7 @@ import ( func TestKeywordSearchTool_Metadata(t *testing.T) { t.Parallel() tool := &KeywordSearchTool{agentID: "test"} - assert.Equal(t, "keyword_search", tool.Name()) + assert.Empty(t, cmp.Diff("keyword_search", tool.Name())) assert.Contains(t, tool.Description(), "FTS5") params := tool.Parameters() require.NotNil(t, params) @@ -30,7 +31,7 @@ func TestKeywordSearchTool_MissingQuery(t *testing.T) { func TestSemanticSearchTool_Metadata(t *testing.T) { t.Parallel() tool := &SemanticSearchTool{agentID: "test"} - assert.Equal(t, "semantic_search", tool.Name()) + assert.Empty(t, cmp.Diff("semantic_search", tool.Name())) assert.Contains(t, tool.Description(), "semantic similarity") params := tool.Parameters() require.NotNil(t, params) @@ -46,7 +47,7 @@ func TestSemanticSearchTool_MissingQuery(t *testing.T) { func TestChunkReadTool_Metadata(t *testing.T) { t.Parallel() tool := &ChunkReadTool{agentID: "test"} - assert.Equal(t, "chunk_read", tool.Name()) + assert.Empty(t, cmp.Diff("chunk_read", tool.Name())) assert.Contains(t, tool.Description(), "full content") params := tool.Parameters() require.NotNil(t, params) diff --git a/pkg/tools/skills_test.go b/pkg/tools/skills_test.go index 36a7e45ff..c8acdd453 100644 --- a/pkg/tools/skills_test.go +++ b/pkg/tools/skills_test.go @@ -6,6 +6,7 @@ import ( "testing" "github.com/ZanzyTHEbar/dragonscale/pkg/skills" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -62,7 +63,7 @@ func TestSkillSearchTool(t *testing.T) { loader := setupTestSkills(t) tool := NewSkillSearchTool(loader) - assert.Equal(t, "skill_search", tool.Name()) + assert.Empty(t, cmp.Diff("skill_search", tool.Name())) t.Run("finds matching skills", func(t *testing.T) { result := tool.Execute(t.Context(), map[string]interface{}{"query": "trading"}) @@ -88,7 +89,7 @@ func TestSkillReadTool(t *testing.T) { loader := setupTestSkills(t) tool := NewSkillReadTool(loader) - assert.Equal(t, "skill_read", tool.Name()) + assert.Empty(t, cmp.Diff("skill_read", tool.Name())) t.Run("reads existing skill", func(t *testing.T) { result := tool.Execute(t.Context(), map[string]interface{}{"name": "risk-management"}) @@ -114,7 +115,7 @@ func TestSkillTraverseTool(t *testing.T) { loader := setupTestSkills(t) tool := NewSkillTraverseTool(loader) - assert.Equal(t, "skill_traverse", tool.Name()) + assert.Empty(t, cmp.Diff("skill_traverse", tool.Name())) t.Run("traverses links at depth 1", func(t *testing.T) { result := tool.Execute(t.Context(), map[string]interface{}{"name": "risk-management"}) diff --git a/pkg/worker/worker_test.go b/pkg/worker/worker_test.go index 91e2e6b0d..5816dfcaa 100644 --- a/pkg/worker/worker_test.go +++ b/pkg/worker/worker_test.go @@ -7,6 +7,7 @@ import ( "testing" "time" + "github.com/google/go-cmp/cmp" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -67,7 +68,7 @@ func TestWorker_RunOnce_HandlerCalled(t *testing.T) { err := worker.RunOnce(ctx, q, opts) require.NoError(t, err) - assert.Equal(t, int32(1), called.Load(), "handler must be called once") + assert.Empty(t, cmp.Diff(int32(1), called.Load()), "handler must be called once") } func TestWorker_RunOnce_JobMarkedSucceeded(t *testing.T) { @@ -90,7 +91,7 @@ func TestWorker_RunOnce_JobMarkedSucceeded(t *testing.T) { done, err := q.GetJob(ctx, sqlc.GetJobParams{ID: job.ID}) require.NoError(t, err) - assert.Equal(t, "succeeded", done.Status) + assert.Empty(t, cmp.Diff("succeeded", done.Status)) } func TestWorker_RunOnce_HandlerError_Requeued(t *testing.T) { @@ -114,8 +115,8 @@ func TestWorker_RunOnce_HandlerError_Requeued(t *testing.T) { requeued, err := q.GetJob(ctx, sqlc.GetJobParams{ID: job.ID}) require.NoError(t, err) - assert.Equal(t, "queued", requeued.Status, "job should be requeued after transient failure") - assert.Equal(t, int64(1), requeued.Attempts, "attempt count must increment") + assert.Empty(t, cmp.Diff("queued", requeued.Status), "job should be requeued after transient failure") + assert.Empty(t, cmp.Diff(int64(1), requeued.Attempts), "attempt count must increment") assert.NotNil(t, requeued.LastError) } @@ -146,7 +147,7 @@ func TestWorker_RunOnce_MaxAttemptsExhausted_MarkedFailed(t *testing.T) { failed++ } } - assert.Equal(t, 1, failed, "job should be permanently failed after exhausting max attempts") + assert.Empty(t, cmp.Diff(1, failed), "job should be permanently failed after exhausting max attempts") } func TestWorker_RunOnce_UnknownKind_MarkedFailed(t *testing.T) { @@ -171,7 +172,7 @@ func TestWorker_RunOnce_UnknownKind_MarkedFailed(t *testing.T) { failed++ } } - assert.Equal(t, 1, failed, "unknown kind should permanently fail the job") + assert.Empty(t, cmp.Diff(1, failed), "unknown kind should permanently fail the job") } func TestWorker_RunOnce_NilQueries(t *testing.T) { @@ -224,5 +225,5 @@ func TestWorker_RunLoop_ProcessesJobs(t *testing.T) { } _ = worker.RunLoop(ctx, q, opts) - assert.Equal(t, int32(total), processed.Load(), "all jobs must be processed") + assert.Empty(t, cmp.Diff(int32(total), processed.Load()), "all jobs must be processed") }