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/mapper.go b/apps/gateway/internal/adapters/providers/openai/mapper.go index ab2c41c..94c2f39 100644 --- a/apps/gateway/internal/adapters/providers/openai/mapper.go +++ b/apps/gateway/internal/adapters/providers/openai/mapper.go @@ -1,6 +1,8 @@ package openai import ( + "strings" + "github.com/dappnode/dappnode-nexus-gateway/pkg/domain" ) @@ -144,7 +146,9 @@ func buildProviderRequest(req domain.GenerateRequest, model domain.PublicModel) if req.Store != nil { body["store"] = *req.Store } - if req.ServiceTier != nil { + if model.ServiceTier != nil && strings.TrimSpace(*model.ServiceTier) != "" { + body["service_tier"] = strings.TrimSpace(*model.ServiceTier) + } else if !strings.EqualFold(strings.TrimSpace(model.ProviderConfig.ProviderName), "doubleword") && req.ServiceTier != nil { body["service_tier"] = *req.ServiceTier } diff --git a/apps/gateway/internal/adapters/providers/openai/policy_test.go b/apps/gateway/internal/adapters/providers/openai/policy_test.go index f4bdaeb..8126e3f 100644 --- a/apps/gateway/internal/adapters/providers/openai/policy_test.go +++ b/apps/gateway/internal/adapters/providers/openai/policy_test.go @@ -59,6 +59,46 @@ func TestBuildProviderRequest_ForwardsReasoningEffort(t *testing.T) { } } +func TestBuildProviderRequest_DoublewordServiceTierComesFromModel(t *testing.T) { + req := domain.GenerateRequest{ + PublicModelID: "doubleword/test-realtime", + Input: []domain.InputItem{message("user", "hi")}, + ServiceTier: stringPtr("priority"), + } + realtime := domain.PublicModel{ + UpstreamModelName: "test-model", + ProviderConfig: domain.ProviderConfig{ProviderName: "doubleword"}, + } + + built := buildProviderRequest(req, realtime) + if _, ok := built.Body["service_tier"]; ok { + t.Fatal("realtime Doubleword model must not forward a caller-selected service tier") + } + + async := realtime + async.ServiceTier = stringPtr(" flex ") + built = buildProviderRequest(req, async) + if got := built.Body["service_tier"]; got != "flex" { + t.Fatalf("service_tier = %v, want configured flex tier", got) + } +} + +func TestBuildProviderRequest_OtherProvidersForwardClientServiceTier(t *testing.T) { + req := domain.GenerateRequest{ + Input: []domain.InputItem{message("user", "hi")}, + ServiceTier: stringPtr("auto"), + } + model := domain.PublicModel{ + UpstreamModelName: "test-model", + ProviderConfig: domain.ProviderConfig{ProviderName: "openai"}, + } + + built := buildProviderRequest(req, model) + if got := built.Body["service_tier"]; got != "auto" { + t.Fatalf("service_tier = %v, want auto", got) + } +} + func TestBuildProviderRequest_DeveloperRolePolicy(t *testing.T) { req := domain.GenerateRequest{ PublicModelID: "test/developer-role", diff --git a/apps/gateway/internal/application/services/generate_service.go b/apps/gateway/internal/application/services/generate_service.go index b5aa472..638deec 100644 --- a/apps/gateway/internal/application/services/generate_service.go +++ b/apps/gateway/internal/application/services/generate_service.go @@ -306,6 +306,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 @@ -484,6 +485,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 95c8572..517fa10 100644 --- a/apps/gateway/internal/application/services/generate_service_test.go +++ b/apps/gateway/internal/application/services/generate_service_test.go @@ -227,6 +227,9 @@ func TestGenerateServiceExecute_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 TestGenerateServiceExecute_ReturnsFallbackFailure(t *testing.T) { @@ -359,6 +362,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}, @@ -369,7 +373,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, @@ -466,6 +470,33 @@ func TestGenerateServiceExecute_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 TestGenerateServiceExecute_DoesNotRouteUnknownModel(t *testing.T) { router := &stubRouterClient{decision: domain.RouteDecision{PublicModelID: "minimax/minimax-m2.7"}} svc := NewGenerateService( 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