package gateway import ( "context" "encoding/json" "errors" "testing" ) type scripted struct { name string err error calls *[]string } func (s *scripted) Complete(ctx context.Context, req Request) (*Response, error) { *s.calls = append(*s.calls, s.name) if s.err != nil { return nil, s.err } return &Response{Text: "answered by " + s.name, Model: s.name}, nil } func gwErr(code string, status int) error { return &Error{Code: code, Status: status, Message: code} } func TestFailoverAsksTheNextProviderOnARateLimit(t *testing.T) { var calls []string f := NewFailover( &scripted{name: "groq", err: gwErr(CodeRateLimited, 429), calls: &calls}, &scripted{name: "cerebras", calls: &calls}, ) resp, err := f.Complete(context.Background(), Request{ Messages: []Message{{Role: RoleUser, Text: "hello"}}, }) if err != nil { t.Fatalf("want an answer from the fallback, got %v", err) } if resp.Model != "cerebras" { t.Errorf("answered by %q, want cerebras", resp.Model) } if len(calls) != 2 || calls[0] != "groq" { t.Errorf("provider order was %v, want groq then cerebras", calls) } } func TestFailoverDoesNotMaskABadCredential(t *testing.T) { // A 401 is THIS deployment's own configuration and fails identically // everywhere. Trying three providers would turn one visible fault into // three invisible ones and leave the operator nothing to fix. var calls []string f := NewFailover( &scripted{name: "groq", err: gwErr(CodeUnauthorized, 401), calls: &calls}, &scripted{name: "cerebras", calls: &calls}, ) _, err := f.Complete(context.Background(), Request{ Messages: []Message{{Role: RoleUser, Text: "hello"}}, }) var e *Error if !errors.As(err, &e) || e.Code != CodeUnauthorized { t.Fatalf("want the unauthorized error raised, got %v", err) } if len(calls) != 1 { t.Errorf("called %v; a terminal error must not reach the fallback", calls) } } func TestFailoverWillNotMoveAConversationBoundToItsProvider(t *testing.T) { // ToolCall.Extra is provider metadata echoed back verbatim — Gemini's // thought signature. Replaying it at a different vendor sends it a field it // cannot read; dropping it kills the vendor that issued it. Either way the // conversation belongs to whoever started it. var calls []string f := NewFailover( &scripted{name: "gemini", err: gwErr(CodeRateLimited, 429), calls: &calls}, &scripted{name: "groq", calls: &calls}, ) _, err := f.Complete(context.Background(), Request{ Messages: []Message{ {Role: RoleUser, Text: "how many open positions?"}, {Role: RoleAssistant, ToolCalls: []ToolCall{{ ID: "c1", Name: "open_positions", Input: json.RawMessage(`{}`), Extra: json.RawMessage(`{"thought_signature":"abc"}`), }}}, }, }) if err == nil { t.Fatal("want the rate limit raised, not a second provider's answer") } if len(calls) != 1 { t.Errorf("called %v; a pinned conversation must not fail over", calls) } } func TestFailoverWithNoFallbacksIsTheProviderItself(t *testing.T) { var calls []string p := &scripted{name: "groq", calls: &calls} if got := NewFailover(p); got != Gateway(p) { t.Error("with no fallbacks the primary must be returned unwrapped") } } func TestFailoverExhaustedReturnsTheLastError(t *testing.T) { var calls []string f := NewFailover( &scripted{name: "a", err: gwErr(CodeRateLimited, 429), calls: &calls}, &scripted{name: "b", err: gwErr(CodeUpstream, 503), calls: &calls}, ) _, err := f.Complete(context.Background(), Request{ Messages: []Message{{Role: RoleUser, Text: "hi"}}, }) var e *Error if !errors.As(err, &e) || e.Status != 503 { t.Fatalf("want the LAST provider's error, got %v", err) } if len(calls) != 2 { t.Errorf("called %v, want both tried", calls) } }