Skip to content
Merged
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
Original file line number Diff line number Diff line change
Expand Up @@ -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?"}]},
Expand All @@ -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" {
Expand Down
2 changes: 2 additions & 0 deletions apps/gateway/internal/adapters/http/mapper/chat_to_domain.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"`
Expand Down Expand Up @@ -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"),
}
Expand Down
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
}
9 changes: 9 additions & 0 deletions apps/gateway/internal/adapters/providers/openai/proxy.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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,
Expand All @@ -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)
}
Expand Down
35 changes: 35 additions & 0 deletions apps/gateway/internal/adapters/providers/openai/proxy_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down Expand Up @@ -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},
Expand All @@ -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,
Expand Down Expand Up @@ -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(
Expand Down
1 change: 1 addition & 0 deletions pkg/domain/generation.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
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