diff --git a/apps/gateway/internal/adapters/http/mapper/chat_mapper_test.go b/apps/gateway/internal/adapters/http/mapper/chat_mapper_test.go index 1dfc9e2..12422af 100644 --- a/apps/gateway/internal/adapters/http/mapper/chat_mapper_test.go +++ b/apps/gateway/internal/adapters/http/mapper/chat_mapper_test.go @@ -14,7 +14,7 @@ func read(t *testing.T, body string) (domain.GenerateRequest, error) { } func TestChatCompletionRequestToDomain_ReadsTheSummary(t *testing.T) { - req, err := read(t, `{"model":"m","stream":true,"max_completion_tokens":50,"parallel_tool_calls":false, + req, err := read(t, `{"model":"m","stream":true,"max_completion_tokens":50,"service_tier":"flex","parallel_tool_calls":false, "response_format":{"type":"json_object"},"temperature":0.2,"unknown_extension":{"x":1}, "messages":[{"role":"developer","content":"be brief"}, {"role":"user","content":[{"type":"text","text":"What is "},{"type":"image_url","image_url":{"url":"data:image/png;base64,AA"}},{"type":"text","text":"this?"}]}, @@ -24,7 +24,7 @@ func TestChatCompletionRequestToDomain_ReadsTheSummary(t *testing.T) { if err != nil { t.Fatal(err) } - if req.PublicModelID != "m" || !req.Stream || *req.MaxOutputTokens != 50 || *req.ParallelToolCalls || !req.StructuredOutput { + if req.PublicModelID != "m" || !req.Stream || *req.MaxOutputTokens != 50 || req.ServiceTier == nil || *req.ServiceTier != "flex" || *req.ParallelToolCalls || !req.StructuredOutput { t.Fatalf("summary %+v", req) } if len(req.Input) != 4 || *req.Input[0].Role != "developer" || *req.Input[1].Content != "What is this?" || req.Input[2].Content != nil || *req.Input[3].Content != "42" { diff --git a/apps/gateway/internal/adapters/http/mapper/chat_to_domain.go b/apps/gateway/internal/adapters/http/mapper/chat_to_domain.go index 3d2c23b..bb71a4a 100644 --- a/apps/gateway/internal/adapters/http/mapper/chat_to_domain.go +++ b/apps/gateway/internal/adapters/http/mapper/chat_to_domain.go @@ -19,6 +19,7 @@ type chatRequest struct { Stream *bool `json:"stream"` MaxTokens *int `json:"max_tokens"` MaxCompletionTokens *int `json:"max_completion_tokens"` + ServiceTier *string `json:"service_tier"` N *int `json:"n"` Tools []chatTool `json:"tools"` ToolChoice json.RawMessage `json:"tool_choice"` @@ -77,6 +78,7 @@ func ChatCompletionRequestToDomain(raw json.RawMessage) (domain.GenerateRequest, PublicModelID: req.Model, Stream: req.Stream != nil && *req.Stream, MaxOutputTokens: req.MaxTokens, + ServiceTier: req.ServiceTier, ParallelToolCalls: req.ParallelToolCalls, StructuredOutput: req.ResponseFormat != nil && (req.ResponseFormat.Type == "json_object" || req.ResponseFormat.Type == "json_schema"), } diff --git a/apps/gateway/internal/adapters/metering/runtime_client.go b/apps/gateway/internal/adapters/metering/runtime_client.go index 8fb2d5f..9a3124f 100644 --- a/apps/gateway/internal/adapters/metering/runtime_client.go +++ b/apps/gateway/internal/adapters/metering/runtime_client.go @@ -330,6 +330,7 @@ func (m runtimePublicModel) toDomain() domain.PublicModel { Description: m.Description, ProviderModelID: m.ProviderModelID, UpstreamModelName: m.UpstreamModelName, + ServiceTier: m.ServiceTier, ProviderConfig: m.ProviderConfig.toDomain(), SupportsChatCompletions: *m.SupportsChatCompletions, SupportsChatCompletionsStream: *m.SupportsChatCompletionsStream, @@ -349,6 +350,7 @@ func (m runtimePublicModel) toDomain() domain.PublicModel { model.Fallback = &domain.ProviderTarget{ ProviderModelID: m.Fallback.ProviderModelID, UpstreamModelName: m.Fallback.UpstreamModelName, + ServiceTier: m.Fallback.ServiceTier, ProviderConfig: m.Fallback.ProviderConfig.toDomain(), } } @@ -418,6 +420,7 @@ type runtimePublicModel struct { Description *string `json:"description,omitempty"` ProviderModelID string `json:"provider_model_id"` UpstreamModelName string `json:"upstream_model_name"` + ServiceTier *string `json:"service_tier,omitempty"` ProviderConfig runtimeProviderConfig `json:"provider_config"` Fallback *runtimeProviderTarget `json:"fallback,omitempty"` SupportsChatCompletions *bool `json:"supports_chat_completions"` @@ -440,6 +443,7 @@ type runtimePublicModel struct { type runtimeProviderTarget struct { ProviderModelID string `json:"provider_model_id"` UpstreamModelName string `json:"upstream_model_name"` + ServiceTier *string `json:"service_tier,omitempty"` ProviderConfig runtimeProviderConfig `json:"provider_config"` } diff --git a/apps/gateway/internal/adapters/metering/runtime_client_test.go b/apps/gateway/internal/adapters/metering/runtime_client_test.go index 84c830f..fe01d0b 100644 --- a/apps/gateway/internal/adapters/metering/runtime_client_test.go +++ b/apps/gateway/internal/adapters/metering/runtime_client_test.go @@ -249,8 +249,10 @@ func TestRuntimeClientRejectsInvalidModelSuccess(t *testing.T) { func TestRuntimeClientMapsProviderFallback(t *testing.T) { response := validRuntimePublicModel("openai/gpt-test") + response.ServiceTier = stringPointer("priority") response.Fallback = &runtimeProviderTarget{ ProviderModelID: "fallback-model", UpstreamModelName: "fallback-upstream", + ServiceTier: stringPointer("flex"), ProviderConfig: runtimeProviderConfig{ ID: "fallback-config", ProviderName: "fallback", BaseURL: "https://fallback.test", APIKeySecretRef: "FALLBACK_API_KEY", Active: true, @@ -267,6 +269,10 @@ func TestRuntimeClientMapsProviderFallback(t *testing.T) { model.Fallback.ProviderConfig.ProviderName != "fallback" { t.Fatalf("fallback = %#v", model.Fallback) } + if model.ServiceTier == nil || *model.ServiceTier != "priority" || + model.Fallback.ServiceTier == nil || *model.Fallback.ServiceTier != "flex" { + t.Fatalf("service tiers = primary %v, fallback %v", model.ServiceTier, model.Fallback.ServiceTier) + } } func TestRuntimeClientModelListRequiresArrayAndValidUniqueItems(t *testing.T) { @@ -566,6 +572,10 @@ func boolPointer(value bool) *bool { return &value } +func stringPointer(value string) *string { + return &value +} + func int64Pointer(value int64) *int64 { return &value } diff --git a/apps/gateway/internal/adapters/providers/openai/proxy.go b/apps/gateway/internal/adapters/providers/openai/proxy.go index 22bfcdb..9555520 100644 --- a/apps/gateway/internal/adapters/providers/openai/proxy.go +++ b/apps/gateway/internal/adapters/providers/openai/proxy.go @@ -6,6 +6,7 @@ import ( "encoding/json" "io" "net/http" + "strings" "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/adapters/observability/metrics" "github.com/dappnode/dappnode-nexus-gateway/apps/gateway/internal/application/ports" @@ -22,6 +23,7 @@ var gatewayOnlyFields = []string{"provider_options"} // provider extensions). The only edits are: // // - model: the provider's name for the public model; +// - service_tier: fixed by the catalog for Doubleword models; // - stream_options.include_usage: on for streams, so usage can be metered; // - the output-token limit: clamped to the model's maximum, and sent as // max_tokens to the providers that only document that name (Novita, @@ -41,6 +43,13 @@ func PrepareProxyBody(raw []byte, model domain.PublicModel, stream bool) (map[st var transforms []string body["model"] = model.UpstreamModelName + // Doubleword's realtime and async catalog entries select the upstream tier. + // Never let a caller change the tier behind a priced public model. + if model.ServiceTier != nil && strings.TrimSpace(*model.ServiceTier) != "" { + body["service_tier"] = strings.TrimSpace(*model.ServiceTier) + } else if strings.EqualFold(model.ProviderConfig.ProviderName, "doubleword") { + delete(body, "service_tier") + } for _, field := range gatewayOnlyFields { delete(body, field) } diff --git a/apps/gateway/internal/adapters/providers/openai/proxy_test.go b/apps/gateway/internal/adapters/providers/openai/proxy_test.go index c759548..631c5ed 100644 --- a/apps/gateway/internal/adapters/providers/openai/proxy_test.go +++ b/apps/gateway/internal/adapters/providers/openai/proxy_test.go @@ -56,6 +56,41 @@ func TestPrepareProxyBody_ForwardsTheClientRequest(t *testing.T) { } } +func TestPrepareProxyBody_DoublewordServiceTier(t *testing.T) { + for _, stream := range []bool{false, true} { + for _, tc := range []struct { + name string + provider string + configured *string + client string + want string + present bool + }{ + {name: "async", provider: "doubleword", configured: stringPtr(" flex "), client: "priority", want: "flex", present: true}, + {name: "realtime", provider: "doubleword", client: "priority"}, + {name: "other provider", provider: "openai", client: "priority", want: "priority", present: true}, + } { + t.Run(tc.name, func(t *testing.T) { + raw := `{"model":"public/model","messages":[],"service_tier":"` + tc.client + `"}` + body, _, err := PrepareProxyBody([]byte(raw), domain.PublicModel{ + UpstreamModelName: "upstream", + ProviderConfig: domain.ProviderConfig{ProviderName: tc.provider}, + ServiceTier: tc.configured, + }, stream) + if err != nil { + t.Fatal(err) + } + got, ok := body["service_tier"] + if ok != tc.present || ok && got != tc.want { + t.Fatalf("service_tier = %v (present %v), want %q (present %v)", got, ok, tc.want, tc.present) + } + }) + } + } +} + +func stringPtr(value string) *string { return &value } + func TestPrepareProxyBody_Edits(t *testing.T) { t.Run("streams always report usage", func(t *testing.T) { got := prepared(t, `{"messages":[],"stream":true,"stream_options":{"include_usage":false,"continuous_usage_stats":true}}`, "novita", true) diff --git a/apps/gateway/internal/application/services/generate_service.go b/apps/gateway/internal/application/services/generate_service.go index 4fc891d..6be0ca9 100644 --- a/apps/gateway/internal/application/services/generate_service.go +++ b/apps/gateway/internal/application/services/generate_service.go @@ -81,6 +81,7 @@ func shouldTryFallback(ctx context.Context, err error, fallback *domain.Provider func withProviderTarget(model domain.PublicModel, target domain.ProviderTarget) domain.PublicModel { model.ProviderModelID = target.ProviderModelID model.UpstreamModelName = target.UpstreamModelName + model.ServiceTier = target.ServiceTier model.ProviderConfig = target.ProviderConfig model.Fallback = nil return model @@ -202,6 +203,10 @@ func (s *GenerateService) validateRequest(endpoint string, req domain.GenerateRe return domain.ErrUnsupportedFeature("structured_output") } + if strings.EqualFold(model.ProviderConfig.ProviderName, "doubleword") && req.ServiceTier != nil { + return domain.ErrInvalidField("service_tier is fixed by the selected Doubleword model") + } + if model.EffectiveProofMode() == domain.ProofModeTinfoilAttestedTransport && !strings.EqualFold(model.ProviderConfig.ProviderName, "tinfoil") { return domain.ErrUnsupportedFeature("Tinfoil verified transport for non-Tinfoil provider") diff --git a/apps/gateway/internal/application/services/generate_service_test.go b/apps/gateway/internal/application/services/generate_service_test.go index 81266cc..70e63ff 100644 --- a/apps/gateway/internal/application/services/generate_service_test.go +++ b/apps/gateway/internal/application/services/generate_service_test.go @@ -239,6 +239,9 @@ func TestProxy_UsesConfiguredFallbackOnce(t *testing.T) { if meter.lastSuccessModel.ProviderConfig.ProviderName != "fallback" { t.Fatalf("metered provider = %q, want fallback", meter.lastSuccessModel.ProviderConfig.ProviderName) } + if meter.lastSuccessModel.ServiceTier == nil || *meter.lastSuccessModel.ServiceTier != "flex" { + t.Fatalf("metered service tier = %v, want flex", meter.lastSuccessModel.ServiceTier) + } } func TestProxy_ReturnsFallbackFailure(t *testing.T) { @@ -299,6 +302,7 @@ func newDirectModelGenerateService(meter *stubUsageMeter, provider *stubProvider } func newFallbackGenerateService(meter *stubUsageMeter, primary, fallback *stubProvider) *GenerateService { + serviceTier := "flex" return NewGenerateService( &stubAuthService{authCtx: domain.AuthContext{ Account: domain.Account{ID: "acc1", Status: domain.AccountStatusActive}, @@ -309,7 +313,7 @@ func newFallbackGenerateService(meter *stubUsageMeter, primary, fallback *stubPr ProviderModelID: "primary-model", UpstreamModelName: "primary-upstream", ProviderConfig: domain.ProviderConfig{ProviderName: "primary"}, - Fallback: &domain.ProviderTarget{ProviderModelID: "fallback-model", UpstreamModelName: "fallback-upstream", ProviderConfig: domain.ProviderConfig{ProviderName: "fallback"}}, + Fallback: &domain.ProviderTarget{ProviderModelID: "fallback-model", UpstreamModelName: "fallback-upstream", ServiceTier: &serviceTier, ProviderConfig: domain.ProviderConfig{ProviderName: "fallback"}}, SupportsChatCompletions: true, SupportsChatCompletionsStream: true, MaxContextWindow: 1000, @@ -377,6 +381,33 @@ func TestProxy_RejectsWhenBalanceIsEmpty(t *testing.T) { } } +func TestValidateRequest_RejectsCallerSelectedDoublewordServiceTier(t *testing.T) { + svc := &GenerateService{} + serviceTier := "flex" + model := domain.PublicModel{ + PublicModelID: "doubleword/test-realtime", + SupportsChatCompletions: true, + ProviderConfig: domain.ProviderConfig{ProviderName: "doubleword"}, + } + err := svc.validateRequest(domain.EndpointChatCompletions, domain.GenerateRequest{ + ServiceTier: &serviceTier, + }, model) + if err == nil { + t.Fatal("expected caller-selected Doubleword service tier to be rejected") + } + gwErr, ok := err.(*domain.GatewayError) + if !ok || gwErr.Code != domain.ErrCodeInvalidField { + t.Fatalf("error = %#v, want invalid_field", err) + } + + model.ProviderConfig.ProviderName = "openai" + if err := svc.validateRequest(domain.EndpointChatCompletions, domain.GenerateRequest{ + ServiceTier: &serviceTier, + }, model); err != nil { + t.Fatalf("non-Doubleword service tier rejected: %v", err) + } +} + func TestProxy_DoesNotRouteUnknownModel(t *testing.T) { router := &stubRouterClient{decision: domain.RouteDecision{PublicModelID: "minimax/minimax-m2.7"}} svc := NewGenerateService( diff --git a/pkg/domain/generation.go b/pkg/domain/generation.go index 84e02d5..c9146e0 100644 --- a/pkg/domain/generation.go +++ b/pkg/domain/generation.go @@ -28,6 +28,7 @@ type GenerateRequest struct { Stream bool MaxOutputTokens *int + ServiceTier *string ParallelToolCalls *bool // StructuredOutput is set when the client asks for JSON output. StructuredOutput bool diff --git a/pkg/domain/model.go b/pkg/domain/model.go index cdc8987..a5c83b6 100644 --- a/pkg/domain/model.go +++ b/pkg/domain/model.go @@ -21,6 +21,7 @@ type ProviderConfig struct { type ProviderTarget struct { ProviderModelID string UpstreamModelName string + ServiceTier *string ProviderConfig ProviderConfig } @@ -32,6 +33,7 @@ type PublicModel struct { Description *string ProviderModelID string UpstreamModelName string + ServiceTier *string ProviderConfig ProviderConfig Fallback *ProviderTarget SupportsChatCompletions bool