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
4 changes: 4 additions & 0 deletions apps/gateway/internal/adapters/metering/runtime_client.go
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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(),
}
}
Expand Down Expand Up @@ -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"`
Expand All @@ -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"`
}

Expand Down
10 changes: 10 additions & 0 deletions apps/gateway/internal/adapters/metering/runtime_client_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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) {
Expand Down Expand Up @@ -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
}
6 changes: 5 additions & 1 deletion apps/gateway/internal/adapters/providers/openai/mapper.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
package openai

import (
"strings"

"github.com/dappnode/dappnode-nexus-gateway/pkg/domain"
)

Expand Down Expand Up @@ -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
}

Expand Down
40 changes: 40 additions & 0 deletions apps/gateway/internal/adapters/providers/openai/policy_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down Expand Up @@ -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},
Expand All @@ -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,
Expand Down Expand Up @@ -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(
Expand Down
2 changes: 2 additions & 0 deletions pkg/domain/model.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ type ProviderConfig struct {
type ProviderTarget struct {
ProviderModelID string
UpstreamModelName string
ServiceTier *string
ProviderConfig ProviderConfig
}

Expand All @@ -32,6 +33,7 @@ type PublicModel struct {
Description *string
ProviderModelID string
UpstreamModelName string
ServiceTier *string
ProviderConfig ProviderConfig
Fallback *ProviderTarget
SupportsChatCompletions bool
Expand Down
Loading