package streamgate import ( "testing" "time" ) // TestNormalizedEventConstructorsAndBaseDisposition verifies that every public // NormalizedEvent constructor produces a value with the correct base disposition // for its kind, passes Validate, and rejects invalid shapes such as empty // channel, zero timestamp, invalid HTTP status, and empty required payloads. func TestNormalizedEventConstructorsAndBaseDisposition(t *testing.T) { ts := time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC) // Table of valid kind → expected disposition, using the public constructor. tests := []struct { name string kind EventKind expectedDisp BaseEventDisposition build func() (NormalizedEvent, error) }{ { name: "response_start", kind: EventKindResponseStart, expectedDisp: BaseDispositionHold, build: func() (NormalizedEvent, error) { return NewResponseStartEvent("ch", 200, map[string]string{"x-custom": "val"}, ts) }, }, { name: "text_delta", kind: EventKindTextDelta, expectedDisp: BaseDispositionReleaseCandidate, build: func() (NormalizedEvent, error) { return NewTextDeltaEvent("ch", "hello", ts) }, }, { name: "reasoning_delta", kind: EventKindReasoningDelta, expectedDisp: BaseDispositionReleaseCandidate, build: func() (NormalizedEvent, error) { return NewReasoningDeltaEvent("ch", "thinking", ts) }, }, { name: "tool_call_fragment", kind: EventKindToolCallFragment, expectedDisp: BaseDispositionReleaseCandidate, build: func() (NormalizedEvent, error) { return NewToolCallFragmentEvent("ch", "call-1", "get_weather", `{"q":"hi"}`, ts) }, }, { name: "terminal", kind: EventKindTerminal, expectedDisp: BaseDispositionTerminalSuccessCandidate, build: func() (NormalizedEvent, error) { return NewTerminalEvent("ch", ts) }, }, { name: "provider_error", kind: EventKindProviderError, expectedDisp: BaseDispositionTerminalErrorCandidate, build: func() (NormalizedEvent, error) { desc, err := NewExternalDescriptor("error.type", "code", "message", "") if err != nil { t.Fatalf("NewExternalDescriptor: %v", err) } return NewProviderErrorEvent("ch", desc, FailureCauseChain{}, ts) }, }, } for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { ev, err := tc.build() if err != nil { t.Fatalf("constructor %s returned error: %v", tc.name, err) } if ev.Kind() != tc.kind { t.Errorf("Kind() = %q, want %q", ev.Kind(), tc.kind) } if ev.Disposition() != tc.expectedDisp { t.Errorf("Disposition() = %q, want %q", ev.Disposition(), tc.expectedDisp) } if err := ev.Validate(); err != nil { t.Errorf("Validate() returned error: %v", err) } }) } // --- Invalid shapes: each constructor must reject bad input --- invalidTests := []struct { name string err func() error }{ {"response_start_empty_channel", func() error { _, err := NewResponseStartEvent("", 200, nil, ts) return err }}, {"response_start_zero_timestamp", func() error { _, err := NewResponseStartEvent("ch", 200, nil, time.Time{}) return err }}, {"response_start_status_below_min", func() error { _, err := NewResponseStartEvent("ch", 99, nil, ts) return err }}, {"response_start_status_above_max", func() error { _, err := NewResponseStartEvent("ch", 600, nil, ts) return err }}, {"text_delta_empty_channel", func() error { _, err := NewTextDeltaEvent("", "hello", ts) return err }}, {"text_delta_empty_payload", func() error { _, err := NewTextDeltaEvent("ch", "", ts) return err }}, {"text_delta_zero_timestamp", func() error { _, err := NewTextDeltaEvent("ch", "hello", time.Time{}) return err }}, {"reasoning_delta_empty_payload", func() error { _, err := NewReasoningDeltaEvent("ch", "", ts) return err }}, {"tool_call_fragment_empty_id", func() error { _, err := NewToolCallFragmentEvent("ch", "", "name", "args", ts) return err }}, {"tool_call_fragment_empty_name", func() error { _, err := NewToolCallFragmentEvent("ch", "id", "", "args", ts) return err }}, {"tool_call_fragment_zero_timestamp", func() error { _, err := NewToolCallFragmentEvent("ch", "id", "name", "args", time.Time{}) return err }}, {"terminal_empty_channel", func() error { _, err := NewTerminalEvent("", ts) return err }}, {"terminal_zero_timestamp", func() error { _, err := NewTerminalEvent("ch", time.Time{}) return err }}, {"provider_error_empty_channel", func() error { desc, err := NewExternalDescriptor("error.type", "code", "message", "") if err != nil { t.Fatalf("NewExternalDescriptor: %v", err) } _, err = NewProviderErrorEvent("", desc, FailureCauseChain{}, ts) return err }}, {"provider_error_zero_timestamp", func() error { desc, err := NewExternalDescriptor("error.type", "code", "message", "") if err != nil { t.Fatalf("NewExternalDescriptor: %v", err) } _, err = NewProviderErrorEvent("ch", desc, FailureCauseChain{}, time.Time{}) return err }}, } for _, tc := range invalidTests { t.Run(tc.name, func(t *testing.T) { if err := tc.err(); err == nil { t.Fatalf("%s: expected error, got nil", tc.name) } }) } // --- Unknown kind: BaseDispositionOf and EventKind.Validate reject it --- t.Run("unknown_kind_rejected", func(t *testing.T) { unknown := EventKind("unknown_kind") if _, err := BaseDispositionOf(unknown); err == nil { t.Fatal("BaseDispositionOf with unknown kind should return error") } if err := unknown.Validate(); err == nil { t.Fatal("EventKind.Validate with unknown kind should return error") } }) } // TestResponseStartRejectsUnsafeHeaders verifies that NewResponseStart rejects // unsafe staged metadata (invalid HTTP status, empty channel, zero timestamp) // and that the headers map is defensively copied so that caller mutation of // the input or accessor return value cannot alter the stored snapshot. func TestResponseStartRejectsUnsafeHeaders(t *testing.T) { ts := time.Date(2025, 1, 1, 0, 0, 0, 0, time.UTC) // --- Invalid HTTP status codes are rejected --- t.Run("status_below_min_rejected", func(t *testing.T) { if _, err := NewResponseStart("ch", 99, nil, ts); err == nil { t.Fatal("status 99 should be rejected") } }) t.Run("status_above_max_rejected", func(t *testing.T) { if _, err := NewResponseStart("ch", 600, nil, ts); err == nil { t.Fatal("status 600 should be rejected") } }) t.Run("status_100_accepted", func(t *testing.T) { rs, err := NewResponseStart("ch", 100, nil, ts) if err != nil { t.Fatalf("status 100 should be accepted: %v", err) } if rs.Status() != 100 { t.Errorf("Status() = %d, want 100", rs.Status()) } }) t.Run("status_599_accepted", func(t *testing.T) { rs, err := NewResponseStart("ch", 599, nil, ts) if err != nil { t.Fatalf("status 599 should be accepted: %v", err) } if rs.Status() != 599 { t.Errorf("Status() = %d, want 599", rs.Status()) } }) // --- Empty channel and zero timestamp are rejected --- t.Run("empty_channel_rejected", func(t *testing.T) { if _, err := NewResponseStart("", 200, nil, ts); err == nil { t.Fatal("empty channel should be rejected") } }) t.Run("zero_timestamp_rejected", func(t *testing.T) { if _, err := NewResponseStart("ch", 200, nil, time.Time{}); err == nil { t.Fatal("zero timestamp should be rejected") } }) // --- Defensive copy: mutating the input map after construction does not // affect the stored headers. --- t.Run("input_mutation_does_not_affect_stored", func(t *testing.T) { input := map[string]string{"x-custom": "val", "x-other": "orig"} rs, err := NewResponseStart("ch", 200, input, ts) if err != nil { t.Fatalf("NewResponseStart: %v", err) } // Mutate the caller-owned input map. input["x-custom"] = "mutated" input["x-new"] = "injected" delete(input, "x-other") got := rs.Headers() if got["x-custom"] != "val" { t.Errorf("x-custom = %q, want %q (input mutation leaked)", got["x-custom"], "val") } if got["x-other"] != "orig" { t.Errorf("x-other = %q, want %q (input deletion leaked)", got["x-other"], "orig") } if _, ok := got["x-new"]; ok { t.Error("x-new should not be present (input injection leaked)") } }) // --- Defensive copy: mutating the accessor return value does not affect // the stored headers. --- t.Run("accessor_mutation_does_not_affect_stored", func(t *testing.T) { rs, err := NewResponseStart("ch", 200, map[string]string{"x-custom": "val"}, ts) if err != nil { t.Fatalf("NewResponseStart: %v", err) } got := rs.Headers() got["x-custom"] = "mutated" got["x-injected"] = "yes" // Second call should return the original values. got2 := rs.Headers() if got2["x-custom"] != "val" { t.Errorf("x-custom = %q, want %q (accessor mutation leaked)", got2["x-custom"], "val") } if _, ok := got2["x-injected"]; ok { t.Error("x-injected should not be present (accessor injection leaked)") } }) // --- Nil headers are accepted; accessor returns a non-nil empty map // (the constructor allocates via make(...)). --- t.Run("nil_headers_accepted", func(t *testing.T) { rs, err := NewResponseStart("ch", 200, nil, ts) if err != nil { t.Fatalf("NewResponseStart with nil headers: %v", err) } got := rs.Headers() if got == nil { t.Fatal("Headers() = nil, want non-nil empty map") } if len(got) != 0 { t.Errorf("Headers() length = %d, want 0", len(got)) } }) // --- NewResponseStartEvent also defensively copies headers. --- t.Run("event_defensive_copy", func(t *testing.T) { input := map[string]string{"x-custom": "val"} ev, err := NewResponseStartEvent("ch", 200, input, ts) if err != nil { t.Fatalf("NewResponseStartEvent: %v", err) } input["x-custom"] = "mutated" rs, err := ev.AsResponseStart() if err != nil { t.Fatalf("AsResponseStart: %v", err) } if rs.Headers()["x-custom"] != "val" { t.Errorf("x-custom = %q, want %q (event input mutation leaked)", rs.Headers()["x-custom"], "val") } }) // --- Forbidden hop-by-hop / transport metadata headers --- forbiddenHeaders := []string{ "Connection", "Keep-Alive", "Proxy-Authenticate", "Proxy-Authorization", "Proxy-Connection", "TE", "Trailer", "Transfer-Encoding", "Upgrade", "Content-Length", } for _, h := range forbiddenHeaders { t.Run("forbidden/"+h, func(t *testing.T) { headers := map[string]string{h: "value"} // NewResponseStart rejects if _, err := NewResponseStart("ch", 200, headers, ts); err == nil { t.Errorf("NewResponseStart should reject forbidden header %q", h) } // NewResponseStartEvent rejects if _, err := NewResponseStartEvent("ch", 200, headers, ts); err == nil { t.Errorf("NewResponseStartEvent should reject forbidden header %q", h) } }) } // --- Mixed-case forbidden headers are rejected case-insensitively --- t.Run("forbidden_mixed_case", func(t *testing.T) { headers := map[string]string{"CONTENT-LENGTH": "value", "Transfer-Encoding": "value"} if _, err := NewResponseStart("ch", 200, headers, ts); err == nil { t.Error("NewResponseStart should reject mixed-case forbidden headers") } if _, err := NewResponseStartEvent("ch", 200, headers, ts); err == nil { t.Error("NewResponseStartEvent should reject mixed-case forbidden headers") } }) // --- Normal extension/end-to-end headers are preserved --- t.Run("normal_headers_preserved", func(t *testing.T) { headers := map[string]string{"x-custom": "val", "etag": "abc", "cache-control": "no-cache"} rs, err := NewResponseStart("ch", 200, headers, ts) if err != nil { t.Fatalf("NewResponseStart: %v", err) } got := rs.Headers() for k, v := range headers { if got[k] != v { t.Errorf("header %q = %q, want %q", k, got[k], v) } } }) // --- Direct Validate on ResponseStart rejects forbidden headers --- t.Run("direct_validate_rejects_forbidden", func(t *testing.T) { rs := ResponseStart{ channel: "ch", status: 200, headers: map[string]string{"Content-Length": "value"}, timestamp: ts, } if err := rs.Validate(); err == nil { t.Error("ResponseStart.Validate should reject forbidden header") } }) // --- Direct Validate on response-start NormalizedEvent rejects forbidden headers --- t.Run("event_validate_rejects_forbidden", func(t *testing.T) { ev := NormalizedEvent{ kind: EventKindResponseStart, channel: "ch", timestamp: ts, disposition: BaseDispositionHold, status: 200, headers: map[string]string{"Transfer-Encoding": "value"}, } if err := ev.Validate(); err == nil { t.Error("NormalizedEvent.Validate should reject forbidden header for response_start") } }) }