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
135 changes: 135 additions & 0 deletions internal/proxy/openai_usage_bypass_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,135 @@
package proxy_test

import (
"context"
"errors"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"

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

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

func oaiBypassBody() []byte {
return []byte(`{"model":"` + bypassRequestedMdl + `","messages":[{"role":"user","content":"hi"}]}`)
}

// TestProxyOpenAI_UsageBypass_BelowThreshold_SkipsScorer: gate on + headroom skips scorer on OpenAI wire.
func TestProxyOpenAI_UsageBypass_BelowThreshold_SkipsScorer(t *testing.T) {
svc, fr, p := bypassFixture(t, 0.20)
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(""))

require.NoError(t, svc.ProxyOpenAIChatCompletion(bypassCtx(0.80), oaiBypassBody(), rec, req))

assert.Equal(t, 0, fr.routeCalls, "scorer must not run while subscription has headroom")
require.Len(t, p.proxyBodies, 1)
assert.Contains(t, string(p.proxyBodies[0]), `"`+bypassRequestedMdl+`"`)
assert.Equal(t, "usage_bypass", rec.Header().Get(proxy.HeaderRouterDecision))
assert.Equal(t, bypassRequestedMdl, rec.Header().Get(proxy.HeaderRouterModel))
assert.Contains(t, rec.Body.String(), `"object":"chat.completion"`)
}

// TestProxyOpenAI_UsageBypass_WeeklyLimit_FallsBackToRoutedDispatch: bypass 429 must reroute via routeFor.
func TestProxyOpenAI_UsageBypass_WeeklyLimit_FallsBackToRoutedDispatch(t *testing.T) {
bypassResp := &providers.UpstreamErrorResponse{
Status: http.StatusTooManyRequests,
Headers: http.Header{
"anthropic-ratelimit-unified-weekly-limit": []string{"100000"},
"anthropic-ratelimit-unified-weekly-reset": []string{"2025-12-31T00:00:00Z"},
"anthropic-ratelimit-unified-weekly-remaining": []string{"0"},
},
Body: []byte(`{"type":"error","error":{"type":"rate_limit_error","message":"weekly limit exceeded"}}`),
}
routedResp := func(w http.ResponseWriter) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte(`{"id":"msg_1","type":"message","role":"assistant","model":"` + bypassScorerPickMdl + `","content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":1}}`))
}
inner := &fakeProvider{proxyResponse: routedResp}
wrappedP := &swapErrProvider{first: bypassResp, second: nil, inner: inner}
fr := &fakeRouter{decision: router.Decision{Provider: providers.ProviderAnthropic, Model: bypassScorerPickMdl, Reason: "cluster:v0.2"}}
obs := usage.NewObserver([]byte("salt"), 10*time.Minute, time.Now)
obs.Record(obs.Key([]byte(bypassSubToken)), usage.Snapshot{
Primary: usage.Window{UsedPercent: 0.20, WindowMinutes: 300},
})
svc := proxy.NewService(fr, map[string]providers.Client{providers.ProviderAnthropic: wrappedP}, nil, false, nil, nil, false, providers.ProviderAnthropic, bypassScorerPickMdl, nil).
WithSubscriptionAwareRouting(obs, 0.05, 2.0)

rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(""))
ctx := bypassCtx(0.80)
ctx = context.WithValue(ctx, proxy.ExternalIDContextKey{}, "org-oai-bypass-reroute")
ctx = context.WithValue(ctx, proxy.InstallationIDContextKey{}, uuid.New().String())

require.NoError(t, svc.ProxyOpenAIChatCompletion(ctx, oaiBypassBody(), rec, req))

assert.Equal(t, 1, fr.routeCalls, "scorer must run once on the reroute after bypass 429")
assert.NotEqual(t, http.StatusTooManyRequests, rec.Code, "the 429 must not be flushed to the client")
assert.Equal(t, bypassScorerPickMdl, rec.Header().Get(proxy.HeaderRouterModel))
assert.Equal(t, "cluster:v0.2", rec.Header().Get(proxy.HeaderRouterDecision))
assert.Contains(t, rec.Body.String(), `"object":"chat.completion"`)
}

// TestSubscriptionOnly_OpenAI_BypassRetryable_Refuses402: subscription-only refuses retryable bypass failure.
func TestSubscriptionOnly_OpenAI_BypassRetryable_Refuses402(t *testing.T) {
bypassResp := &providers.UpstreamErrorResponse{
Status: http.StatusTooManyRequests,
Body: []byte(`{"type":"error","error":{"type":"rate_limit_error","message":"weekly limit exceeded"}}`),
}
p := &fakeProvider{proxyErr: bypassResp}
wrappedP := &swapErrProvider{first: bypassResp, second: nil, inner: p}
fr := &fakeRouter{decision: router.Decision{Provider: providers.ProviderAnthropic, Model: bypassScorerPickMdl, Reason: "cluster:v0.2"}}
obs := usage.NewObserver([]byte("salt"), 10*time.Minute, time.Now)
obs.Record(obs.Key([]byte(bypassSubToken)), usage.Snapshot{
Primary: usage.Window{UsedPercent: 0.20, WindowMinutes: 300},
})
svc := proxy.NewService(fr, map[string]providers.Client{providers.ProviderAnthropic: wrappedP}, nil, false, nil, nil, false, providers.ProviderAnthropic, bypassScorerPickMdl, nil).
WithSubscriptionAwareRouting(obs, 0.05, 2.0)

rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(""))
err := svc.ProxyOpenAIChatCompletion(billing.WithSubscriptionOnly(bypassCtx(0.80)), oaiBypassBody(), rec, req)
require.Error(t, err)
assert.True(t, errors.Is(err, proxy.ErrCreditsExhaustedSubscriptionUnavailable))
assert.Equal(t, 0, fr.routeCalls)
assert.Equal(t, 1, wrappedP.calls)
}

// TestSubscriptionOnly_OpenAI_BypassRetryable_Stream_NoPartialCommit: streaming Prelude stays buffered until Discard on 402.
func TestSubscriptionOnly_OpenAI_BypassRetryable_Stream_NoPartialCommit(t *testing.T) {
bypassResp := &providers.UpstreamErrorResponse{
Status: http.StatusTooManyRequests,
Body: []byte(`{"type":"error","error":{"type":"rate_limit_error","message":"weekly limit exceeded"}}`),
}
p := &fakeProvider{proxyErr: bypassResp}
wrappedP := &swapErrProvider{first: bypassResp, second: nil, inner: p}
fr := &fakeRouter{decision: router.Decision{Provider: providers.ProviderAnthropic, Model: bypassScorerPickMdl, Reason: "cluster:v0.2"}}
obs := usage.NewObserver([]byte("salt"), 10*time.Minute, time.Now)
obs.Record(obs.Key([]byte(bypassSubToken)), usage.Snapshot{
Primary: usage.Window{UsedPercent: 0.20, WindowMinutes: 300},
})
svc := proxy.NewService(fr, map[string]providers.Client{providers.ProviderAnthropic: wrappedP}, nil, false, nil, nil, false, providers.ProviderAnthropic, bypassScorerPickMdl, nil).
WithSubscriptionAwareRouting(obs, 0.05, 2.0)

rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(""))
body := []byte(`{"model":"` + bypassRequestedMdl + `","stream":true,"messages":[{"role":"user","content":"hi"}]}`)
err := svc.ProxyOpenAIChatCompletion(billing.WithSubscriptionOnly(bypassCtx(0.80)), body, rec, req)
require.Error(t, err)
assert.True(t, errors.Is(err, proxy.ErrCreditsExhaustedSubscriptionUnavailable))
assert.Equal(t, 0, fr.routeCalls)
assert.Equal(t, 1, wrappedP.calls)
assert.Empty(t, rec.Body.String(), "retryable bypass must Discard buffered Prelude; client must see no SSE bytes")
assert.False(t, rec.Flushed, "HTTP status must not be committed before the 402 mapping")
}
35 changes: 33 additions & 2 deletions internal/proxy/service.go
Original file line number Diff line number Diff line change
Expand Up @@ -4186,11 +4186,42 @@ func (s *Service) ProxyOpenAIChatCompletion(ctx context.Context, body []byte, w
}
routeStart := time.Now()
routeRes, err := s.runTurnLoop(ctx, env, feats, apiKeyID, installationID, subAgentHint, r.Header, routeRequest)
routeMs := time.Since(routeStart).Milliseconds()
if err != nil {
log.Error("Routing failed for OpenAI request", "err", err, "route_ms", routeMs, "requested_model", feats.Model, "total_input_tokens", feats.Tokens)
log.Error("Routing failed for OpenAI request", "err", err, "route_ms", time.Since(routeStart).Milliseconds(), "requested_model", feats.Model, "total_input_tokens", feats.Tokens)
return err
}

// Anthropic UsageBypass consumer — same contract as ProxyMessages.
if routeRes.UsageBypass && routeRes.Decision.Provider == providers.ProviderAnthropic {
err := s.bypassAnthropicOpenAI(ctx, env, feats, routeRes.modelSwitched(), requestStart, requestID, externalID, r, w)
if !errors.Is(err, errBypassRetryable) {
s.firePolicyShadowForServingDecision(ctx, routeRes.Decision, routeRequest)
return err
}
if billing.SubscriptionOnlyFromContext(ctx) {
log.Info("Subscription-only OpenAI bypass hit retryable error; refusing instead of paid reroute",
"request_id", requestID, "external_id", externalID)
return ErrCreditsExhaustedSubscriptionUnavailable
}
routeRequest.SubsidizedModelCostFactor = s.subsidyFactors(ctx, r.Header)
if s.pinStore != nil {
role := roleForTier(catalog.TierFor(feats.Model))
pin, _ := s.loadPin(ctx, sessionKey, role)
hmmHistory := s.loadHMMHistory(ctx, sessionKey, role)
routeRes.SessionKey = sessionKey
routeRes.PriorServedModel, routeRes.SessionEverSwitched = switchHistoryFromPins(pin, hmmHistory)
}
routeRes.UsageBypass = false
decision, rerouteErr := s.routeFor(ctx, routeRequest)
if rerouteErr != nil {
log.Error("Reroute after OpenAI usage-bypass failure failed", "err", rerouteErr)
return rerouteErr
}
routeRes.Decision = decision
routeRes.Fresh = decision
}

routeMs := time.Since(routeStart).Milliseconds()
routeRes.SuggestionMode = r.Header.Get("x-weave-suggestion-mode") == "true"
decision := routeRes.Decision
s.firePolicyShadowForServingDecision(ctx, decision, routeRequest)
Expand Down
4 changes: 2 additions & 2 deletions internal/proxy/turnloop.go
Original file line number Diff line number Diff line change
Expand Up @@ -57,8 +57,8 @@ type turnLoopResult struct {
StickyHit bool
HardPinned bool
// UsageBypass is true when the caller's own subscription has headroom:
// ProxyMessages must serve the requested model straight through with no
// billing debit, bypassing Decision's normal dispatch.
// ProxyMessages / ProxyOpenAIChatCompletion must serve the requested model
// straight through with no billing debit, bypassing Decision's normal dispatch.
UsageBypass bool
PinTier string
PinAgeSec int64
Expand Down
135 changes: 135 additions & 0 deletions internal/proxy/usage_bypass.go
Original file line number Diff line number Diff line change
Expand Up @@ -367,3 +367,138 @@ func (s *Service) bypassToAnthropic(
)
return proxyErr
}

// bypassAnthropicOpenAI is OpenAI-wire bypassToAnthropic: Anthropic upstream, OpenAI response, skip billing.
func (s *Service) bypassAnthropicOpenAI(
ctx context.Context,
env *translate.RequestEnvelope,
feats translate.RoutingFeatures,
modelSwitched bool,
requestStart time.Time,
requestID, externalID string,
r *http.Request,
w http.ResponseWriter,
) error {
log := observability.FromContext(ctx)
decision := router.Decision{
Provider: providers.ProviderAnthropic,
Model: feats.Model,
Reason: "usage_bypass",
}
w.Header().Set(HeaderRouterDecision, decision.Reason)
w.Header().Set(HeaderRouterProvider, decision.Provider)
w.Header().Set(HeaderRouterModel, decision.Model)

p, provErr := s.provider(providers.ProviderAnthropic)
if provErr != nil {
return provErr
}

ctx = resolveAndInjectCredentials(ctx, decision.Provider, r.Header)

outputReserve := contextWindowOutputReserve
if feats.MaxTokens > outputReserve {
outputReserve = feats.MaxTokens
}
opts := translate.EmitOptions{
TargetModel: decision.Model,
TargetProvider: decision.Provider,
Capabilities: router.Lookup(decision.Model),
IncludeStreamUsage: s.usageRequired(),
EnableExtendedContext: shouldEnableExtendedContext(env.FullTokenEstimate(), outputReserve),
ModelSwitched: modelSwitched,
}
prep, emitErr := env.PrepareAnthropic(r.Header, opts)
if emitErr != nil {
log.Error("Failed to emit Anthropic body on OpenAI usage-bypass path", "err", emitErr)
return fmt.Errorf("emit bypass body: %w", emitErr)
}

// Buffer synthetic preamble (subscription-only marker) so a retryable bypass
// failure can Discard and return 402 without a partial SSE stream on the wire.
preludeBuf := newPreludeBuffer(w)
sink := http.ResponseWriter(preludeBuf)
if billing.SubscriptionOnlyFromContext(ctx) {
mw := translate.NewOpenAIRoutingMarkerWriter(preludeBuf, decision.Model, subscriptionOnlyWarningMarker)
if err := mw.Prelude(env.Stream()); err != nil {
log.Error("OpenAI usage-bypass routing-marker prelude failed", "err", err)
}
sink = mw
Comment thread
cursor[bot] marked this conversation as resolved.
}

var extractor *otel.UsageExtractor
var usageSink otel.UsageSink
if s.usageRequired() {
extractor = otel.NewUsageExtractor(nil, providers.ProviderAnthropic)
usageSink = extractor
}
translator := translate.NewSSETranslator(sink, decision.Model, usageSink)
preludeBuf.Seal()

proxyStart := time.Now()
proxyErr := p.Proxy(ctx, decision, prep, translator, r)
if providers.IsRetryable(proxyErr) {
if !preludeBuf.Committed() {
preludeBuf.Discard()
}
return errBypassRetryable
}
proxyErr = finalizeAfterProxy(proxyErr, translator.Finalize)

var upstreamErr *providers.UpstreamErrorResponse
if errors.As(proxyErr, &upstreamErr) {
if !preludeBuf.Committed() {
preludeBuf.Discard()
flushBufferedIfPresent(w, proxyErr)
} else if env.Stream() {
_ = emitOpenAISSEErrorEvent(sink, proxyErr)
} else {
flushBufferedIfPresent(w, proxyErr)
}
proxyErr = nil
}

in, out := extractor.Tokens()
cacheCreation, cacheRead := extractor.CacheTokens()
pricing, _ := catalog.PriceFor(decision.Provider, decision.Model)
inputCost := catalog.EffectiveInputCost(in, cacheCreation, cacheRead, pricing.InputUSDPer1M, pricing, decision.Provider)
outputCost := catalog.EffectiveOutputCost(out, pricing.OutputUSDPer1M)

clientID := ClientIdentityFrom(ctx)
otel.Record(ctx, otel.Span{
Name: "router.usage_bypass",
Start: requestStart,
End: time.Now(),
Attrs: otel.NewAttrBuilder(18).
String("request_id", requestID).
String("external_id", externalID).
String("router_user_id", auth.UserIDFrom(ctx)).
String("client.app", clientID.ClientApp).
String("client.session_id", clientID.SessionID).
String("requested.model", decision.Model).
String("decision.model", decision.Model).
String("decision.provider", decision.Provider).
String("decision.reason", decision.Reason).
Bool("cost.subscription_served", servedOnSubscription(ctx)).
Int64("usage.input_tokens", int64(in)).
Int64("usage.output_tokens", int64(out)).
Int64("usage.cache_creation_input_tokens", int64(cacheCreation)).
Int64("usage.cache_read_input_tokens", int64(cacheRead)).
Float64("cost.requested_input_usd", inputCost).
Float64("cost.requested_output_usd", outputCost).
Float64("cost.actual_input_usd", inputCost).
Float64("cost.actual_output_usd", outputCost).
Build(),
})
otel.Flush(ctx)
log.Info("ProxyOpenAIChatCompletion usage-bypass complete",
"request_id", requestID,
"external_id", externalID,
"requested_model", feats.Model,
"decision_model", decision.Model,
"proxy_ms", time.Since(proxyStart).Milliseconds(),
"total_ms", time.Since(requestStart).Milliseconds(),
"proxy_err", proxyErr,
)
return proxyErr
}