Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
215 changes: 215 additions & 0 deletions internal/proxy/openai_baseline_failover_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,215 @@
package proxy_test

import (
"context"
"io"
"net/http"
"net/http/httptest"
"strings"
"sync"
"testing"

"workweave/router/internal/providers"
"workweave/router/internal/providers/anthropic"
"workweave/router/internal/providers/openaicompat"
"workweave/router/internal/proxy"
"workweave/router/internal/router"
"workweave/router/internal/translate"

"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)

// anthropicMessageJSON is a minimal Anthropic Messages body for OpenAI→Anthropic translation.
const anthropicMessageJSON = `{"id":"msg_1","type":"message","role":"assistant","model":"claude-opus-4-8","content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn","usage":{"input_tokens":5,"output_tokens":1}}`

// TestProxyOpenAI_OSSOutageFailsOverToBaselineAnthropic: OSS outage → Anthropic baseline on the OpenAI wire.
func TestProxyOpenAI_OSSOutageFailsOverToBaselineAnthropic(t *testing.T) {
var (
mu sync.Mutex
fireworksCount int
openRouterCount int
anthropicCount int
anthropicReceivedModel string
)

fail503 := func(counter *int) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
mu.Lock()
*counter++
mu.Unlock()
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusServiceUnavailable)
_, _ = w.Write([]byte(`{"error":{"message":"provider unavailable"}}`))
}
}
fireworks := httptest.NewServer(fail503(&fireworksCount))
defer fireworks.Close()
openrouter := httptest.NewServer(fail503(&openRouterCount))
defer openrouter.Close()

anthropicUpstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
body, _ := io.ReadAll(r.Body)
mu.Lock()
anthropicCount++
anthropicReceivedModel = gjson.GetBytes(body, "model").String()
mu.Unlock()
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(anthropicMessageJSON))
}))
defer anthropicUpstream.Close()

store := newFakePinStore()
tel := newCaptureTelemetry()
svc := proxy.NewService(
&fakeRouter{decision: router.Decision{Provider: "fireworks", Model: "deepseek/deepseek-v4-pro"}},
map[string]providers.Client{
"fireworks": openaicompat.NewClient("test-fw-key", fireworks.URL),
"openrouter": openaicompat.NewClient("test-or-key", openrouter.URL),
"anthropic": anthropic.NewClient("test-anthropic-key", anthropicUpstream.URL),
},
nil, false, nil, store, false, providers.ProviderAnthropic, "claude-haiku-4-5", tel,
).WithDeploymentKeyedProviders(map[string]struct{}{
"fireworks": {},
"openrouter": {},
"anthropic": {},
})

rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(""))
body := []byte(`{"model":"claude-opus-4-8","messages":[{"role":"user","content":"hi"}]}`)

err := svc.ProxyOpenAIChatCompletion(authedCtx("11111111-1111-1111-1111-111111111111"), body, rec, req)
require.NoError(t, err, "ProxyOpenAIChatCompletion should succeed via baseline failover to Anthropic")

mu.Lock()
defer mu.Unlock()
assert.Equal(t, 1, fireworksCount, "Fireworks (primary OSS binding) tried once")
assert.Equal(t, 1, openRouterCount, "OpenRouter (OSS fallback binding) tried once")
assert.Equal(t, 1, anthropicCount, "Anthropic baseline failover dispatched once")
assert.Equal(t, "claude-opus-4-8", anthropicReceivedModel, "baseline failover must request the caller's model on Anthropic")

assert.Equal(t, http.StatusOK, rec.Code)
assert.Equal(t, "anthropic", rec.Header().Get(proxy.HeaderRouterProvider), "served provider header reflects the baseline failover")
assert.Equal(t, "claude-opus-4-8", rec.Header().Get(proxy.HeaderRouterModel), "x-router-model reflects the baseline model that served")
assert.Contains(t, rec.Body.String(), `"object":"chat.completion"`)
assert.NotContains(t, rec.Body.String(), "deepseek/deepseek-v4-pro")

require.NotEmpty(t, store.usages, "baseline failover must write pin usage")
assert.Equal(t, "claude-opus-4-8", store.usages[len(store.usages)-1].ServedModel, "pin usage records the served baseline model")
}

// TestProxyOpenAI_ForcedModelUnavailableDoesNotSubstituteAnthropic: forced-model must not substitute Anthropic.
func TestProxyOpenAI_ForcedModelUnavailableDoesNotSubstituteAnthropic(t *testing.T) {
var anthropicCount int
var googleCount int
anthropicProv := &fakeProvider{proxyResponse: func(w http.ResponseWriter) {
anthropicCount++
w.WriteHeader(http.StatusOK)
}}
google := &fakeProvider{proxyResponse: func(w http.ResponseWriter) {
googleCount++
w.WriteHeader(http.StatusOK)
}}
svc := makeProxyService(
router.Decision{
Provider: providers.ProviderGoogle,
Model: "gemini-3.1-pro-preview",
Reason: translate.ReasonUserForceModel,
},
map[string]providers.Client{
providers.ProviderAnthropic: anthropicProv,
providers.ProviderGoogle: google,
},
).WithDeploymentKeyedProviders(map[string]struct{}{providers.ProviderAnthropic: {}})

rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(""))
body := []byte(`{"model":"claude-opus-4-8","messages":[{"role":"user","content":"hi"}]}`)

err := svc.ProxyOpenAIChatCompletion(context.Background(), body, rec, req)
require.Error(t, err)
assert.Equal(t, 0, anthropicCount, "forced model requests must never substitute Anthropic")
assert.Equal(t, 0, googleCount, "unwired forced provider must not be dispatched")
}

// TestProxyOpenAI_OSSOutageNoBaselineWhenRequestedModelIsOSS: no baseline when the caller requested OSS.
func TestProxyOpenAI_OSSOutageNoBaselineWhenRequestedModelIsOSS(t *testing.T) {
var anthropicCount int
anthropicUpstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
anthropicCount++
w.WriteHeader(http.StatusOK)
}))
defer anthropicUpstream.Close()

fail := func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusServiceUnavailable)
_, _ = w.Write([]byte(`{"error":{"message":"down"}}`))
}
fireworks := httptest.NewServer(http.HandlerFunc(fail))
defer fireworks.Close()
openrouter := httptest.NewServer(http.HandlerFunc(fail))
defer openrouter.Close()

svc := proxy.NewService(
&fakeRouter{decision: router.Decision{Provider: "fireworks", Model: "deepseek/deepseek-v4-pro"}},
map[string]providers.Client{
"fireworks": openaicompat.NewClient("k", fireworks.URL),
"openrouter": openaicompat.NewClient("k", openrouter.URL),
"anthropic": anthropic.NewClient("k", anthropicUpstream.URL),
},
nil, false, nil, nil, false, providers.ProviderAnthropic, "claude-haiku-4-5", nil,
).WithDeploymentKeyedProviders(map[string]struct{}{
"fireworks": {}, "openrouter": {}, "anthropic": {},
})

rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(""))
body := []byte(`{"model":"deepseek/deepseek-v4-pro","messages":[{"role":"user","content":"hi"}]}`)

_ = svc.ProxyOpenAIChatCompletion(context.Background(), body, rec, req)
assert.Equal(t, 0, anthropicCount, "baseline failover must not fire when the caller requested the OSS model")
}

// TestProxyOpenAI_OSSOutageNoBaselineWhenAnthropicExcluded: no baseline when Anthropic is excluded.
func TestProxyOpenAI_OSSOutageNoBaselineWhenAnthropicExcluded(t *testing.T) {
var anthropicCount int
anthropicUpstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
anthropicCount++
w.WriteHeader(http.StatusOK)
}))
defer anthropicUpstream.Close()

fail := func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusServiceUnavailable)
_, _ = w.Write([]byte(`{"error":{"message":"provider unavailable"}}`))
}
fireworks := httptest.NewServer(http.HandlerFunc(fail))
defer fireworks.Close()
openrouter := httptest.NewServer(http.HandlerFunc(fail))
defer openrouter.Close()

svc := proxy.NewService(
&fakeRouter{decision: router.Decision{Provider: "fireworks", Model: "deepseek/deepseek-v4-pro"}},
map[string]providers.Client{
"fireworks": openaicompat.NewClient("k", fireworks.URL),
"openrouter": openaicompat.NewClient("k", openrouter.URL),
"anthropic": anthropic.NewClient("k", anthropicUpstream.URL),
},
nil, false, nil, nil, false, providers.ProviderAnthropic, "claude-haiku-4-5", nil,
).WithDeploymentKeyedProviders(map[string]struct{}{
"fireworks": {}, "openrouter": {}, "anthropic": {},
}).WithExcludedProvidersOverride([]string{providers.ProviderAnthropic})

rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(""))
body := []byte(`{"model":"claude-opus-4-8","messages":[{"role":"user","content":"hi"}]}`)

_ = svc.ProxyOpenAIChatCompletion(context.Background(), body, rec, req)
assert.Equal(t, 0, anthropicCount, "baseline failover must not hit Anthropic when it is excluded")
assert.NotEqual(t, providers.ProviderAnthropic, rec.Header().Get(proxy.HeaderRouterProvider), "served provider must not be the excluded Anthropic")
}
163 changes: 163 additions & 0 deletions internal/proxy/openai_subscription_failover_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,163 @@
package proxy_test

import (
"context"
"fmt"
"net/http"
"net/http/httptest"
"strings"
"sync"
"testing"
"time"

"workweave/router/internal/providers"
"workweave/router/internal/proxy"
"workweave/router/internal/proxy/usage"
"workweave/router/internal/router"

"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

const (
oaiSubFailoverToken = "sk-ant-oat01-openai-sub-failover-token"
oaiSubFailoverModel = "claude-haiku-4-5"
oaiSubFailoverOK = `{"id":"msg_1","type":"message","role":"assistant","content":[{"type":"text","text":"hi"}],"model":"claude-haiku-4-5","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":1}}`
)

// oauthRejectThenDeployOK fails under OAuth subscription creds, then succeeds on the deploy-key path.
type oauthRejectThenDeployOK struct {
mu sync.Mutex
onOAuth error
calls []oauthCallSnap
}

type oauthCallSnap struct {
nilCreds bool
oauth bool
source string
key string
err error
}

func (p *oauthRejectThenDeployOK) Proxy(ctx context.Context, decision router.Decision, prep providers.PreparedRequest, w http.ResponseWriter, r *http.Request) error {
creds := proxy.CredentialsFromContext(ctx)
snap := oauthCallSnap{nilCreds: creds == nil}
if creds != nil {
snap.oauth = creds.OAuth
snap.source = creds.Source
snap.key = string(creds.APIKey)
}
p.mu.Lock()
defer p.mu.Unlock()
if creds != nil && creds.OAuth {
snap.err = p.onOAuth
p.calls = append(p.calls, snap)
return p.onOAuth
}
p.calls = append(p.calls, snap)
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(oaiSubFailoverOK))
return nil
}

func (p *oauthRejectThenDeployOK) Passthrough(context.Context, providers.PreparedRequest, http.ResponseWriter, *http.Request) error {
return fmt.Errorf("unused")
}

func oaiSubSlackObs(t *testing.T) *usage.Observer {
t.Helper()
obs := usage.NewObserver([]byte("salt"), 10*time.Minute, time.Now)
// Slack: preemptive suppress must not fire; only reactive subscription retry recovers.
obs.Record(obs.Key([]byte(oaiSubFailoverToken)), usage.Snapshot{
Primary: usage.Window{UsedPercent: 0.50, WindowMinutes: 300},
})
return obs
}

func oaiSubCtx() context.Context {
return context.WithValue(context.Background(), proxy.AnthropicSubscriptionContextKey{}, oaiSubFailoverToken)
}

func oaiSubMainLoopBody() []byte {
return []byte(`{"model":"` + oaiSubFailoverModel + `","max_tokens":4096,"messages":[{"role":"user","content":"Refactor the auth middleware and add tests."}],"tools":[{"type":"function","function":{"name":"edit_file","parameters":{"type":"object","properties":{}}}}]}`)
}

func oaiSubFailoverSvc(t *testing.T, p providers.Client) *proxy.Service {
t.Helper()
fr := &fakeRouter{decision: router.Decision{Provider: providers.ProviderAnthropic, Model: oaiSubFailoverModel, Reason: "cluster:test"}}
return proxy.NewService(fr, map[string]providers.Client{providers.ProviderAnthropic: p}, nil, false, nil, nil, false, providers.ProviderAnthropic, oaiSubFailoverModel, nil).
WithSubscriptionAwareRouting(oaiSubSlackObs(t), 0.05, 2.0).
WithDeploymentKeyedProviders(map[string]struct{}{providers.ProviderAnthropic: {}})
}

// TestProxyOpenAI_SubscriptionRetry_Live429FailsOverToDeployKey: live 429 on OAuth → Weave/BYOK retry.
func TestProxyOpenAI_SubscriptionRetry_Live429FailsOverToDeployKey(t *testing.T) {
reject := &providers.UpstreamErrorResponse{
Status: http.StatusTooManyRequests,
Body: []byte(`{"type":"error","error":{"type":"rate_limit_error","message":"weekly limit exceeded"}}`),
}
p := &oauthRejectThenDeployOK{onOAuth: reject}
svc := oaiSubFailoverSvc(t, p)

rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(""))
err := svc.ProxyOpenAIChatCompletion(oaiSubCtx(), oaiSubMainLoopBody(), rec, req)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)

p.mu.Lock()
defer p.mu.Unlock()
require.GreaterOrEqual(t, len(p.calls), 2, "must retry after the subscription 429")
assert.True(t, p.calls[0].oauth && p.calls[0].source == "subscription")
last := p.calls[len(p.calls)-1]
assert.True(t, last.nilCreds || !last.oauth, "final dispatch must use the deploy key, not the spent OAuth token")
assert.Contains(t, rec.Body.String(), `"object":"chat.completion"`)
}

// TestProxyOpenAI_SubscriptionRetry_OAuth401FailsOverToDeployKey: OAuth authentication_error → deploy key.
func TestProxyOpenAI_SubscriptionRetry_OAuth401FailsOverToDeployKey(t *testing.T) {
reject := &providers.UpstreamErrorResponse{
Status: http.StatusUnauthorized,
Body: []byte(`{"type":"error","error":{"type":"authentication_error","message":"Invalid authentication credentials"}}`),
}
p := &oauthRejectThenDeployOK{onOAuth: reject}
svc := oaiSubFailoverSvc(t, p)

rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(""))
err := svc.ProxyOpenAIChatCompletion(oaiSubCtx(), oaiSubMainLoopBody(), rec, req)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)

p.mu.Lock()
defer p.mu.Unlock()
require.GreaterOrEqual(t, len(p.calls), 2)
assert.True(t, p.calls[0].oauth)
last := p.calls[len(p.calls)-1]
assert.True(t, last.nilCreds || !last.oauth)
}

// TestProxyOpenAI_SubscriptionRetry_OAuth403FailsOverToDeployKey: OAuth permission_error → deploy key.
func TestProxyOpenAI_SubscriptionRetry_OAuth403FailsOverToDeployKey(t *testing.T) {
reject := &providers.UpstreamErrorResponse{
Status: http.StatusForbidden,
Body: []byte(`{"type":"error","error":{"type":"permission_error","message":"OAuth authentication is currently not allowed for this organization."}}`),
}
p := &oauthRejectThenDeployOK{onOAuth: reject}
svc := oaiSubFailoverSvc(t, p)

rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(""))
err := svc.ProxyOpenAIChatCompletion(oaiSubCtx(), oaiSubMainLoopBody(), rec, req)
require.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)

p.mu.Lock()
defer p.mu.Unlock()
require.GreaterOrEqual(t, len(p.calls), 2)
assert.True(t, p.calls[0].oauth)
last := p.calls[len(p.calls)-1]
assert.True(t, last.nilCreds || !last.oauth)
}
Loading