From 07a3a4a739fb75e62bb2b02ad4e3f352b77c67ba Mon Sep 17 00:00:00 2001 From: David Gageot Date: Wed, 7 Oct 2026 11:03:00 +0200 Subject: [PATCH] feat: price OpenAI fast/priority and ultrafast by actual service tier Cost accounting now uses the tier the provider actually served instead of the one requested, covering chat, responses, and websocket transports. Custom endpoints, routers, and cost overrides are left alone. WebSocket and HTTP paths now resolve the base URL the same way, fixed at client construction for both transports. Assisted-By: docker-agent --- docs/providers/openai/index.md | 8 +- examples/openai-service-tier.yaml | 11 +- pkg/chat/chat.go | 4 +- pkg/chat/chat_test.go | 18 ++ pkg/model/provider/oaistream/adapter.go | 5 + pkg/model/provider/oaistream/adapter_test.go | 49 ++++ pkg/model/provider/openai/client.go | 9 +- pkg/model/provider/openai/response_stream.go | 6 + .../openai/response_stream_terminal_test.go | 58 +++++ .../provider/openai/service_tier_test.go | 239 ++++++++++++++++++ pkg/modelsdev/service_tier.go | 55 ++++ pkg/modelsdev/service_tier_test.go | 80 ++++++ pkg/runtime/cost_test.go | 87 +++++++ pkg/runtime/loop.go | 21 +- pkg/runtime/native_compaction.go | 2 +- pkg/runtime/service_tier_cost_test.go | 104 ++++++++ 16 files changed, 745 insertions(+), 11 deletions(-) create mode 100644 pkg/modelsdev/service_tier.go create mode 100644 pkg/modelsdev/service_tier_test.go create mode 100644 pkg/runtime/service_tier_cost_test.go diff --git a/docs/providers/openai/index.md b/docs/providers/openai/index.md index 7c05723b8..805e5a7c4 100644 --- a/docs/providers/openai/index.md +++ b/docs/providers/openai/index.md @@ -57,7 +57,7 @@ Starting with GPT-5.6, OpenAI renamed the `-mini`/`-nano` size tiers to `-terra` Find more model names at [modelnames.ai](https://modelnames.ai/) or in the [official OpenAI docs](https://platform.openai.com/docs/models). -## Service Tier (Fast Mode) +## Service Tier (Fast and Ultrafast Modes) Set `provider_opts.service_tier` to request OpenAI's [Fast mode](https://developers.openai.com/api/docs/guides/fast-mode): @@ -74,8 +74,12 @@ OpenAI also accepts `priority` for Fast mode. It provides faster processing at p The value is forwarded unchanged to Chat Completions (including reranking) and Responses requests, over either SSE or WebSocket. OpenAI-compatible providers using these APIs also receive the option when set; the endpoint must support it. Other tiers, such as `auto`, `default`, and `flex`, can also be requested; availability and valid values depend on the API and model. When omitted or empty, no `service_tier` is sent, leaving the API's default behavior unchanged. Non-string values are ignored. +For GPT-6 Astra, request [Ultrafast mode](https://developers.openai.com/api/docs/guides/ultrafast-mode) with `service_tier: ultrafast`. + +Cost estimates use the **actual response tier**, not the requested tier. For OpenAI models with no custom `base_url`, Fast (`fast` or `priority`) applies 2× catalogue rates to `gpt-5.6` (the Sol alias), `gpt-5.6-sol`, `gpt-5.6-terra`, `gpt-5.6-luna`, `gpt-6-astra`, `gpt-6-sol`, `gpt-6-luna`, and `gpt-6.1-sol`. Ultrafast applies 6× rates to `gpt-6-astra` only. These adjustments cover input, cached input, cache writes, and output, including the applicable long-context band. A response reporting `default` retains standard pricing, even if Fast was requested. + > [!WARNING] -> Docker Agent's cost estimates do not automatically adjust for `service_tier`. By default, they use catalogue pricing, which can underestimate premium-tier charges. Set the model's [`cost` override](../../configuration/models/index.md#custom-token-pricing) to the applicable input, output, and cache token rates for your tier. +> Missing response tiers, unlisted models or tiers, custom endpoints (including `OPENAI_BASE_URL`), rule-based routers, and other providers retain catalogue pricing. Gateway model IDs already priced as `-fast` are not multiplied again. Estimates do not include endpoint surcharges or negotiated discounts. Use a model's [`cost` override](../../configuration/models/index.md#custom-token-pricing) when automatic pricing does not apply; overrides replace the entire price table and are never multiplied by the service tier. See [`examples/openai-service-tier.yaml`](https://github.com/docker/docker-agent/blob/main/examples/openai-service-tier.yaml) for a complete example. diff --git a/examples/openai-service-tier.yaml b/examples/openai-service-tier.yaml index 985c188c5..3b34ab650 100644 --- a/examples/openai-service-tier.yaml +++ b/examples/openai-service-tier.yaml @@ -1,10 +1,10 @@ -# Fast mode uses premium pricing on supported OpenAI models. -# Set models.fast-gpt.cost to your tier's rates for accurate cost estimates. +# Cost estimates use the actual response tier for supported OpenAI models. # https://developers.openai.com/api/docs/guides/fast-mode +# https://developers.openai.com/api/docs/guides/ultrafast-mode agents: root: - model: fast-gpt + model: fast-gpt # Switch to ultrafast-astra to use Ultrafast mode. description: An assistant using OpenAI Fast mode. instruction: You are a helpful assistant. @@ -14,3 +14,8 @@ models: model: gpt-5.6 provider_opts: service_tier: fast # priority is also accepted by OpenAI + ultrafast-astra: + provider: openai + model: gpt-6-astra + provider_opts: + service_tier: ultrafast diff --git a/pkg/chat/chat.go b/pkg/chat/chat.go index 8b5751d0d..c7aa8a8d6 100644 --- a/pkg/chat/chat.go +++ b/pkg/chat/chat.go @@ -216,6 +216,8 @@ type Usage struct { CachedInputTokens int64 `json:"cached_input_tokens"` CacheWriteTokens int64 `json:"cached_write_tokens"` ReasoningTokens int64 `json:"reasoning_tokens,omitempty"` + // ServiceTier is the actual processing tier reported by the provider, not the requested tier. + ServiceTier string `json:"service_tier,omitempty"` } // PromptTokens sums the disjoint fresh, cache-read, and cache-write input buckets. @@ -223,7 +225,7 @@ func (u *Usage) PromptTokens() int64 { return u.InputTokens + u.CachedInputTokens + u.CacheWriteTokens } -// Add accumulates other's token counts into u. A nil other is a no-op so +// Add accumulates other's token counts into u, not per-request service tiers. A nil other is a no-op so // callers can pass a message's optional usage without checking. func (u *Usage) Add(other *Usage) { if other == nil { diff --git a/pkg/chat/chat_test.go b/pkg/chat/chat_test.go index d4379bad3..79965d232 100644 --- a/pkg/chat/chat_test.go +++ b/pkg/chat/chat_test.go @@ -1,6 +1,7 @@ package chat import ( + "encoding/json" "fmt" "os" "path/filepath" @@ -217,3 +218,20 @@ func TestUsagePromptTokens(t *testing.T) { assert.Equal(t, int64(35), u.PromptTokens(), "prompt is the sum of the three input buckets, excluding output") assert.Zero(t, (&Usage{}).PromptTokens()) } + +func TestUsageServiceTier(t *testing.T) { + t.Parallel() + + usage := &Usage{InputTokens: 10, ServiceTier: "ultrafast"} + wire, err := json.Marshal(usage) + require.NoError(t, err) + var decoded Usage + require.NoError(t, json.Unmarshal(wire, &decoded)) + assert.Equal(t, *usage, decoded) + + var total Usage + total.Add(usage) + total.Add(&Usage{InputTokens: 20, ServiceTier: "default"}) + assert.Equal(t, int64(30), total.InputTokens) + assert.Empty(t, total.ServiceTier, "aggregate usage has no single processing tier") +} diff --git a/pkg/model/provider/oaistream/adapter.go b/pkg/model/provider/oaistream/adapter.go index fed748ff3..3450a9bf8 100644 --- a/pkg/model/provider/oaistream/adapter.go +++ b/pkg/model/provider/oaistream/adapter.go @@ -21,6 +21,7 @@ type StreamAdapter struct { lastFinishReason chat.FinishReason toolCalls map[int]string trackUsage bool + serviceTier string } func NewStreamAdapter(stream *ssestream.Stream[openai.ChatCompletionChunk], trackUsage bool) *StreamAdapter { @@ -42,6 +43,9 @@ func (a *StreamAdapter) Recv() (chat.MessageStreamResponse, error) { } openaiResponse := a.stream.Current() + if openaiResponse.JSON.ServiceTier.Valid() { + a.serviceTier = string(openaiResponse.ServiceTier) + } // Convert the OpenAI response to our generic format response := chat.MessageStreamResponse{ @@ -144,6 +148,7 @@ func (a *StreamAdapter) Recv() (chat.MessageStreamResponse, error) { response.Usage = &chat.Usage{ InputTokens: usage.PromptTokens, OutputTokens: usage.CompletionTokens, + ServiceTier: a.serviceTier, } if usage.JSON.PromptTokensDetails.Valid() { // chat.Usage treats InputTokens, CachedInputTokens and diff --git a/pkg/model/provider/oaistream/adapter_test.go b/pkg/model/provider/oaistream/adapter_test.go index 0f185f35b..998db0f81 100644 --- a/pkg/model/provider/oaistream/adapter_test.go +++ b/pkg/model/provider/oaistream/adapter_test.go @@ -220,3 +220,52 @@ data: [DONE] assert.Equal(t, "Hi", resp.Choices[0].Delta.Content) assert.Empty(t, resp.Choices[0].Delta.ReasoningContent) } + +func TestStreamAdapter_ServiceTier(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name, early, final, want string + }{ + {name: "fast", final: `"fast"`, want: "fast"}, + {name: "priority alias", final: `"priority"`, want: "priority"}, + {name: "ultrafast", final: `"ultrafast"`, want: "ultrafast"}, + {name: "early tier retained", early: `"fast"`, want: "fast"}, + {name: "terminal downgrade wins", early: `"fast"`, final: `"default"`, want: "default"}, + {name: "null retains earlier tier", early: `"ultrafast"`, final: `null`, want: "ultrafast"}, + {name: "missing tier"}, + {name: "null tier", final: `null`}, + {name: "unknown tier preserved", final: `"future-tier"`, want: "future-tier"}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + early, final := "", "" + if tc.early != "" { + early = `,"service_tier":` + tc.early + } + if tc.final != "" { + final = `,"service_tier":` + tc.final + } + sse := `data: {"id":"c1","object":"chat.completion.chunk","created":1,"model":"gpt-6-astra","choices":[{"index":0,"delta":{"content":"Hi"}}]` + early + "}\n\n" + + `data: {"id":"c1","object":"chat.completion.chunk","created":1,"model":"gpt-6-astra","choices":[],"usage":{"prompt_tokens":100,"completion_tokens":7,"total_tokens":107,"prompt_tokens_details":{"cached_tokens":20,"cache_write_tokens":30},"completion_tokens_details":{"reasoning_tokens":3}}` + final + "}\n\ndata: [DONE]\n\n" + for _, tracking := range []bool{true, false} { + adapter := NewStreamAdapter(newTestStream(t, sse), tracking) + t.Cleanup(adapter.Close) + _, err := adapter.Recv() + require.NoError(t, err) + resp, err := adapter.Recv() + require.NoError(t, err) + if !tracking { + assert.Nil(t, resp.Usage) + continue + } + require.NotNil(t, resp.Usage) + assert.Equal(t, tc.want, resp.Usage.ServiceTier) + assert.Equal(t, int64(50), resp.Usage.InputTokens) + assert.Equal(t, int64(100), resp.Usage.PromptTokens()) + assert.Equal(t, int64(7), resp.Usage.OutputTokens) + assert.Equal(t, int64(3), resp.Usage.ReasoningTokens) + } + }) + } +} diff --git a/pkg/model/provider/openai/client.go b/pkg/model/provider/openai/client.go index ad04de0bd..bf1c29259 100644 --- a/pkg/model/provider/openai/client.go +++ b/pkg/model/provider/openai/client.go @@ -9,6 +9,7 @@ import ( "fmt" "log/slog" "net/http" + "os" "slices" "strings" "sync" @@ -60,6 +61,10 @@ func NewClient(ctx context.Context, cfg *latest.ModelConfig, env environment.Pro } globalOptions := options.Apply(opts...) + resolvedBaseURL := "https://api.openai.com/v1" + if globalOptions.Gateway() == "" { + resolvedBaseURL = cmp.Or(cfg.BaseURL, os.Getenv("OPENAI_BASE_URL"), resolvedBaseURL) + } var clientFn func(context.Context) (*openai.Client, error) if gateway := globalOptions.Gateway(); gateway == "" { @@ -190,6 +195,7 @@ func NewClient(ctx context.Context, cfg *latest.ModelConfig, env environment.Pro ModelConfig: *cfg, ModelOptions: globalOptions, Env: env, + BaseURL: resolvedBaseURL, }, clientFn: clientFn, } @@ -198,8 +204,7 @@ func NewClient(ctx context.Context, cfg *latest.ModelConfig, env environment.Pro // The pool is cheap (no connections opened until the first Stream call) // and eager init avoids a data race on the lazy path. if webSocketEnabled(cfg, &globalOptions) { - baseURL := cmp.Or(cfg.BaseURL, "https://api.openai.com/v1") - client.wsPool = newWSPool(httpToWSURL(baseURL), client.buildWSHeaderFn()) + client.wsPool = newWSPool(httpToWSURL(resolvedBaseURL), client.buildWSHeaderFn()) } return client, nil diff --git a/pkg/model/provider/openai/response_stream.go b/pkg/model/provider/openai/response_stream.go index 34f1ba4e5..c8514339a 100644 --- a/pkg/model/provider/openai/response_stream.go +++ b/pkg/model/provider/openai/response_stream.go @@ -24,6 +24,7 @@ var _ responseEventStream = (*ssestream.Stream[responses.ResponseStreamEventUnio type ResponseStreamAdapter struct { stream responseEventStream trackUsage bool + serviceTier string responseState *chat.OpenAIResponse preserveOutput bool responseItems map[int64]json.RawMessage @@ -96,6 +97,9 @@ func (a *ResponseStreamAdapter) Recv() (chat.MessageStreamResponse, error) { } event := a.stream.Current() + if event.Response.JSON.ServiceTier.Valid() { + a.serviceTier = string(event.Response.ServiceTier) + } slog.Debug("Stream event received", "type", event.Type) response := chat.MessageStreamResponse{} @@ -378,6 +382,7 @@ func (a *ResponseStreamAdapter) Recv() (chat.MessageStreamResponse, error) { CachedInputTokens: u.InputTokensDetails.CachedTokens, CacheWriteTokens: u.InputTokensDetails.CacheWriteTokens, ReasoningTokens: u.OutputTokensDetails.ReasoningTokens, + ServiceTier: a.serviceTier, } } // Check if there were any tool calls in the output @@ -420,6 +425,7 @@ func (a *ResponseStreamAdapter) Recv() (chat.MessageStreamResponse, error) { CachedInputTokens: u.InputTokensDetails.CachedTokens, CacheWriteTokens: u.InputTokensDetails.CacheWriteTokens, ReasoningTokens: u.OutputTokensDetails.ReasoningTokens, + ServiceTier: a.serviceTier, } } finishReason := chat.FinishReasonLength diff --git a/pkg/model/provider/openai/response_stream_terminal_test.go b/pkg/model/provider/openai/response_stream_terminal_test.go index 4de70b722..e32916111 100644 --- a/pkg/model/provider/openai/response_stream_terminal_test.go +++ b/pkg/model/provider/openai/response_stream_terminal_test.go @@ -85,3 +85,61 @@ func TestResponseStream_FailedReturnsError(t *testing.T) { assert.Contains(t, err.Error(), "upstream exploded") assert.Contains(t, err.Error(), "resp_789") } + +func TestResponseStream_ServiceTier(t *testing.T) { + t.Parallel() + + for _, terminal := range []string{"response.completed", "response.done", "response.incomplete"} { + for _, tc := range []struct { + name, early, final, want string + }{ + {name: "fast", final: "fast", want: "fast"}, + {name: "priority alias", final: "priority", want: "priority"}, + {name: "ultrafast", final: "ultrafast", want: "ultrafast"}, + {name: "early tier retained", early: "fast", want: "fast"}, + {name: "terminal downgrade wins", early: "fast", final: "default", want: "default"}, + {name: "null retains earlier tier", early: "ultrafast", final: "null", want: "ultrafast"}, + {name: "missing tier"}, + {name: "null tier", final: "null"}, + {name: "unknown tier", final: "future-tier", want: "future-tier"}, + } { + t.Run(terminal+"/"+tc.name, func(t *testing.T) { + t.Parallel() + early := map[string]any{"id": "resp_tier"} + if tc.early != "" { + early["service_tier"] = tc.early + } + final := map[string]any{ + "id": "resp_tier", + "usage": map[string]any{ + "input_tokens": 100, "output_tokens": 7, "total_tokens": 107, + "input_tokens_details": map[string]any{"cached_tokens": 20, "cache_write_tokens": 30}, + "output_tokens_details": map[string]any{"reasoning_tokens": 3}, + }, + } + if tc.final != "" { + final["service_tier"] = tc.final + if tc.final == "null" { + final["service_tier"] = nil + } + } + events := decodeEvents(t, []map[string]any{ + {"type": "response.created", "response": early}, + {"type": terminal, "response": final}, + }) + adapter := newResponseStreamAdapter(&fakeEventStream{events: events}, true) + defer adapter.Close() + _, err := adapter.Recv() + require.NoError(t, err) + resp, err := adapter.Recv() + require.NoError(t, err) + require.NotNil(t, resp.Usage) + assert.Equal(t, tc.want, resp.Usage.ServiceTier) + assert.Equal(t, int64(50), resp.Usage.InputTokens) + assert.Equal(t, int64(100), resp.Usage.PromptTokens()) + assert.Equal(t, int64(7), resp.Usage.OutputTokens) + assert.Equal(t, int64(3), resp.Usage.ReasoningTokens) + }) + } + } +} diff --git a/pkg/model/provider/openai/service_tier_test.go b/pkg/model/provider/openai/service_tier_test.go index b6bf16816..e8fa047f5 100644 --- a/pkg/model/provider/openai/service_tier_test.go +++ b/pkg/model/provider/openai/service_tier_test.go @@ -15,6 +15,7 @@ import ( "github.com/docker/docker-agent/pkg/chat" "github.com/docker/docker-agent/pkg/config/latest" "github.com/docker/docker-agent/pkg/environment" + "github.com/docker/docker-agent/pkg/model/provider/options" "github.com/docker/docker-agent/pkg/rag/types" ) @@ -145,3 +146,241 @@ func TestServiceTier(t *testing.T) { }) } } + +func TestServiceTierUsageTransports(t *testing.T) { + t.Parallel() + + for _, api := range []string{"chat", "responses", "websocket"} { + for _, tc := range []struct { + requested, actual string + }{ + {requested: "fast", actual: "priority"}, + {requested: "fast", actual: "fast"}, + {requested: "ultrafast", actual: "ultrafast"}, + {requested: "fast", actual: "default"}, + {requested: "ultrafast", actual: "default"}, + {actual: "fast"}, + {requested: "fast"}, + } { + t.Run(api+"/"+tc.requested+"/"+tc.actual, func(t *testing.T) { + t.Parallel() + completed := map[string]any{ + "type": "response.completed", + "response": map[string]any{ + "id": "resp_tier", "service_tier": tc.actual, + "usage": map[string]any{ + "input_tokens": 100, "output_tokens": 7, "total_tokens": 107, + "input_tokens_details": map[string]any{"cached_tokens": 20, "cache_write_tokens": 30}, + "output_tokens_details": map[string]any{"reasoning_tokens": 3}, + }, + }, + } + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var payload map[string]any + if api == "websocket" { + conn, err := (&websocket.Upgrader{}).Upgrade(w, r, nil) + if !assert.NoError(t, err) { + return + } + defer conn.Close() + if !assert.NoError(t, conn.ReadJSON(&payload)) { + return + } + if tc.requested == "" { + assert.NotContains(t, payload, "service_tier") + } else { + assert.Equal(t, tc.requested, payload["service_tier"]) + } + assert.NoError(t, conn.WriteJSON(completed)) + return + } + if !assert.NoError(t, json.NewDecoder(r.Body).Decode(&payload)) { + return + } + if tc.requested == "" { + assert.NotContains(t, payload, "service_tier") + } else { + assert.Equal(t, tc.requested, payload["service_tier"]) + } + w.Header().Set("Content-Type", "text/event-stream") + var event any = completed + if api == "chat" { + event = map[string]any{ + "id": "c1", "object": "chat.completion.chunk", "created": 1, "model": "gpt-6-astra", + "service_tier": tc.actual, "choices": []any{}, + "usage": map[string]any{ + "prompt_tokens": 100, "completion_tokens": 7, "total_tokens": 107, + "prompt_tokens_details": map[string]any{"cached_tokens": 20, "cache_write_tokens": 30}, + "completion_tokens_details": map[string]any{"reasoning_tokens": 3}, + }, + } + } + data, err := json.Marshal(event) + if !assert.NoError(t, err) { + return + } + _, _ = io.WriteString(w, "data: "+string(data)+"\n\ndata: [DONE]\n\n") + })) + defer server.Close() + opts := map[string]any{} + if tc.requested != "" { + opts["service_tier"] = tc.requested + } + if api == "chat" { + opts["api_type"] = "openai_chatcompletions" + } + if api == "websocket" { + opts["transport"] = "websocket" + } + client, err := NewClient(t.Context(), &latest.ModelConfig{ + Provider: "openai", Model: "gpt-6-astra", BaseURL: server.URL, TokenKey: "MY_TOKEN", ProviderOpts: opts, + }, environment.NewMapEnvProvider(map[string]string{"MY_TOKEN": "secret"})) + require.NoError(t, err) + defer client.Close() + stream, err := client.CreateChatCompletionStream(t.Context(), []chat.Message{{Role: chat.MessageRoleUser, Content: "hi"}}, nil) + require.NoError(t, err) + defer stream.Close() + var usage *chat.Usage + for { + resp, err := stream.Recv() + if err != nil { + require.ErrorIs(t, err, io.EOF) + break + } + if resp.Usage != nil { + usage = resp.Usage + } + } + require.NotNil(t, usage) + assert.Equal(t, tc.actual, usage.ServiceTier) + assert.Equal(t, int64(50), usage.InputTokens) + assert.Equal(t, int64(100), usage.PromptTokens()) + assert.Equal(t, int64(7), usage.OutputTokens) + assert.Equal(t, int64(3), usage.ReasoningTokens) + }) + } + } +} + +func TestNewClientPricingEndpoint(t *testing.T) { + t.Setenv("OPENAI_BASE_URL", "https://example.com/v1") + env := environment.NewMapEnvProvider(map[string]string{"OPENAI_API_KEY": "secret"}) + client, err := NewClient(t.Context(), &latest.ModelConfig{Provider: "openai", Model: "gpt-6-astra"}, env) + require.NoError(t, err) + defer client.Close() + assert.Equal(t, "https://example.com/v1", client.BaseConfig().BaseURL) + + // Accounting uses the endpoint captured at construction, not the current environment. + t.Setenv("OPENAI_BASE_URL", "https://other.example/v1") + assert.Equal(t, "https://example.com/v1", client.BaseConfig().BaseURL) + + direct, err := NewClient(t.Context(), &latest.ModelConfig{Provider: "openai", Model: "gpt-6-astra", BaseURL: "https://explicit.example/v1"}, env) + require.NoError(t, err) + defer direct.Close() + assert.Equal(t, "https://explicit.example/v1", direct.BaseConfig().BaseURL) + + gateway, err := NewClient(t.Context(), &latest.ModelConfig{Provider: "openai", Model: "gpt-6-astra"}, env, options.WithGateway("https://gateway.example")) + require.NoError(t, err) + defer gateway.Close() + assert.Equal(t, "https://api.openai.com/v1", gateway.BaseConfig().BaseURL) +} + +func TestNewClientPricingEndpointTransports(t *testing.T) { + for _, tc := range []struct { + name, api string + websocket bool + want []string + }{ + {name: "chat", api: "openai_chatcompletions", want: []string{"chat"}}, + {name: "websocket", api: "openai_responses", websocket: true, want: []string{"websocket"}}, + {name: "websocket to SSE", api: "openai_responses", want: []string{"websocket", "sse"}}, + } { + t.Run(tc.name, func(t *testing.T) { + requests := make(chan string, 3) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "Bearer secret", r.Header.Get("Authorization")) + if r.Header.Get("Upgrade") != "" { + requests <- "websocket" + assert.Equal(t, "/v1/responses", r.URL.Path) + if !tc.websocket { + w.WriteHeader(http.StatusNotFound) + return + } + conn, err := (&websocket.Upgrader{}).Upgrade(w, r, nil) + if !assert.NoError(t, err) { + return + } + defer conn.Close() + var payload map[string]any + if !assert.NoError(t, conn.ReadJSON(&payload)) { + return + } + assert.Equal(t, "response.create", payload["type"]) + assert.Equal(t, "ultrafast", payload["service_tier"]) + event := completedEvent("resp_endpoint") + event["response"].(map[string]any)["service_tier"] = "ultrafast" + assert.NoError(t, conn.WriteJSON(event)) + return + } + if tc.api == "openai_chatcompletions" { + requests <- "chat" + assert.Equal(t, "/v1/chat/completions", r.URL.Path) + writeSSEResponse(w) + } else { + requests <- "sse" + assert.Equal(t, "/v1/responses", r.URL.Path) + w.Header().Set("Content-Type", "text/event-stream") + event := completedEvent("resp_endpoint") + event["response"].(map[string]any)["service_tier"] = "ultrafast" + data, err := json.Marshal(event) + if !assert.NoError(t, err) { + return + } + _, _ = io.WriteString(w, "data: "+string(data)+"\n\ndata: [DONE]\n\n") + } + })) + defer server.Close() + baseURL := server.URL + "/v1" + t.Setenv("OPENAI_BASE_URL", baseURL) + client, err := NewClient(t.Context(), &latest.ModelConfig{ + Provider: "openai", Model: "gpt-6-astra", TokenKey: "OPENAI_API_KEY", ProviderOpts: map[string]any{"transport": "websocket", "api_type": tc.api, "service_tier": "ultrafast"}, + }, environment.NewMapEnvProvider(map[string]string{"OPENAI_API_KEY": "secret"})) + require.NoError(t, err) + defer client.Close() + assert.Equal(t, baseURL, client.BaseConfig().BaseURL) + require.NotNil(t, client.wsPool) + assert.Equal(t, httpToWSURL(baseURL), client.wsPool.wsURL) + // Endpoint selection is fixed at construction for both transports. + t.Setenv("OPENAI_BASE_URL", "https://changed.example/v1") + stream, err := client.CreateChatCompletionStream(t.Context(), []chat.Message{{Role: chat.MessageRoleUser, Content: "hi"}}, nil) + require.NoError(t, err) + defer stream.Close() + var usage *chat.Usage + for { + resp, err := stream.Recv() + if err != nil { + require.ErrorIs(t, err, io.EOF) + break + } + if resp.Usage != nil { + usage = resp.Usage + } + } + require.NotNil(t, usage) + if tc.api == "openai_responses" { + assert.Equal(t, "ultrafast", usage.ServiceTier) + } + var got []string + for { + select { + case request := <-requests: + got = append(got, request) + default: + assert.Equal(t, tc.want, got) + assert.Equal(t, baseURL, client.BaseConfig().BaseURL) + return + } + } + }) + } +} diff --git a/pkg/modelsdev/service_tier.go b/pkg/modelsdev/service_tier.go new file mode 100644 index 000000000..a27c68896 --- /dev/null +++ b/pkg/modelsdev/service_tier.go @@ -0,0 +1,55 @@ +package modelsdev + +import ( + "slices" + "time" +) + +// ForServiceTier adjusts catalogue rates for the actual response tier without mutating c. +// Unlisted models, providers and tiers retain their catalogue pricing. +func (c *Cost) ForServiceTier(id ID, tier string) *Cost { + if c == nil || id.Provider != "openai" { + return c + } + model := id.Model + if len(model) > len("-2006-01-02") { + cut := len(model) - len("2006-01-02") + if model[cut-1] == '-' { + if _, err := time.Parse(time.DateOnly, model[cut:]); err == nil { + model = model[:cut-1] + } + } + } + + // https://developers.openai.com/api/docs/pricing; don't extrapolate to future models. + factor := 1.0 + switch tier { + case "fast", "priority": + switch model { + case "gpt-5.6", "gpt-5.6-sol", "gpt-5.6-terra", "gpt-5.6-luna", + "gpt-6-astra", "gpt-6-sol", "gpt-6-luna", "gpt-6.1-sol": + factor = 2 + } + case "ultrafast": + if model == "gpt-6-astra" { + factor = 6 + } + } + if factor == 1 { + return c + } + + out := *c + out.Input *= factor + out.Output *= factor + out.CacheRead *= factor + out.CacheWrite *= factor + out.Tiers = slices.Clone(c.Tiers) + for i := range out.Tiers { + out.Tiers[i].Input *= factor + out.Tiers[i].Output *= factor + out.Tiers[i].CacheRead *= factor + out.Tiers[i].CacheWrite *= factor + } + return &out +} diff --git a/pkg/modelsdev/service_tier_test.go b/pkg/modelsdev/service_tier_test.go new file mode 100644 index 000000000..ee04a2ca4 --- /dev/null +++ b/pkg/modelsdev/service_tier_test.go @@ -0,0 +1,80 @@ +package modelsdev + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestCostForServiceTier(t *testing.T) { + t.Parallel() + + cost := &Cost{ + Input: 10, Output: 50, CacheRead: 1, CacheWrite: 12.5, + Tiers: []CostTier{{ + Rates: Rates{Input: 20, Output: 75, CacheRead: 2, CacheWrite: 25}, + Tier: TierSpec{Type: "context", Size: 272_000}, + }}, + } + for _, tier := range []string{"fast", "priority", "ultrafast"} { + t.Run(tier, func(t *testing.T) { + t.Parallel() + factor := 2.0 + if tier == "ultrafast" { + factor = 6 + } + got := cost.ForServiceTier(NewID("openai", "gpt-6-astra"), tier) + require.NotSame(t, cost, got) + assert.Equal(t, Rates{Input: 10 * factor, Output: 50 * factor, CacheRead: factor, CacheWrite: 12.5 * factor}, got.RatesFor(272_000)) + assert.Equal(t, Rates{Input: 20 * factor, Output: 75 * factor, CacheRead: 2 * factor, CacheWrite: 25 * factor}, got.RatesFor(272_001)) + assert.Equal(t, cost.Tiers[0].Tier, got.Tiers[0].Tier) + assert.Equal(t, Rates{Input: 10, Output: 50, CacheRead: 1, CacheWrite: 12.5}, cost.RatesFor(272_000)) + assert.Equal(t, Rates{Input: 20, Output: 75, CacheRead: 2, CacheWrite: 25}, cost.RatesFor(272_001), "shared catalogue tiers must not be mutated") + }) + } +} + +func TestCostForServiceTierModels(t *testing.T) { + t.Parallel() + + cost := &Cost{Input: 1} + for _, model := range []string{"gpt-5.6", "gpt-5.6-sol", "gpt-5.6-terra", "gpt-5.6-luna", "gpt-6-astra", "gpt-6-sol", "gpt-6-luna", "gpt-6.1-sol"} { + for _, tier := range []string{"fast", "priority"} { + assert.InDelta(t, 2.0, cost.ForServiceTier(NewID("openai", model), tier).Input, 1e-9, "%s/%s", model, tier) + assert.InDelta(t, 2.0, cost.ForServiceTier(NewID("openai", model+"-2026-09-04"), tier).Input, 1e-9, "dated %s/%s", model, tier) + } + } + assert.InDelta(t, 6.0, cost.ForServiceTier(NewID("openai", "gpt-6-astra-2026-09-04"), "ultrafast").Input, 1e-9) +} + +func TestCostForServiceTierUnchanged(t *testing.T) { + t.Parallel() + + cost := &Cost{Input: 1} + for _, tc := range []struct { + provider, model, tier string + }{ + {"openai", "gpt-6-astra", ""}, + {"openai", "gpt-6-astra", "default"}, + {"openai", "gpt-6-astra", "auto"}, + {"openai", "gpt-6-astra", "flex"}, + {"openai", "gpt-6-astra", "scale"}, + {"openai", "gpt-6-astra", "future-tier"}, + {"openai", "gpt-5.6-luna", "ultrafast"}, + {"openai", "gpt-6.1-sol", "ultrafast"}, + {"openai", "gpt-5.5", "fast"}, + {"openai", "gpt-7-astra", "fast"}, + {"openai", "gpt-6-astra-fast", "fast"}, + {"openai", "gpt-6-astra-extra", "fast"}, + {"openai", "gpt-6-astra-2026-02-30", "fast"}, + {"vercel", "openai/gpt-6-astra-fast", "fast"}, + {"azure", "gpt-6-astra", "fast"}, + {"chatgpt", "gpt-6-astra", "ultrafast"}, + {"custom", "gpt-6-astra", "fast"}, + } { + assert.Same(t, cost, cost.ForServiceTier(NewID(tc.provider, tc.model), tc.tier), "%+v", tc) + } + var missing *Cost + assert.Nil(t, missing.ForServiceTier(NewID("openai", "gpt-6-astra"), "ultrafast")) +} diff --git a/pkg/runtime/cost_test.go b/pkg/runtime/cost_test.go index 6bbc8807f..68fccdba7 100644 --- a/pkg/runtime/cost_test.go +++ b/pkg/runtime/cost_test.go @@ -1,6 +1,7 @@ package runtime import ( + "slices" "testing" "github.com/stretchr/testify/assert" @@ -8,6 +9,7 @@ import ( "github.com/docker/docker-agent/pkg/chat" "github.com/docker/docker-agent/pkg/config/latest" + "github.com/docker/docker-agent/pkg/model/provider/base" "github.com/docker/docker-agent/pkg/modelsdev" ) @@ -140,3 +142,88 @@ func TestConfigCostReplacesContextTiers(t *testing.T) { require.NotNil(t, got) assert.InDelta(t, 4.0, *got, 1e-9, "overrides must not mutate shared catalog tiers") } + +func TestApplyModelCostServiceTier(t *testing.T) { + t.Parallel() + + store := modelsdev.NewDatabaseStore(modelsdev.EmbeddedSnapshot()) + for _, tc := range []struct { + name, provider, model, tier string + cfg latest.ModelConfig + short, long float64 + }{ + {name: "Luna fast", provider: "openai", model: "gpt-5.6-luna", tier: "fast", short: 0.0478, long: 0.4496}, + {name: "Luna priority alias", provider: "openai", model: "gpt-5.6-luna", tier: "priority", short: 0.0478, long: 0.4496}, + {name: "Astra fast", provider: "openai", model: "gpt-6-astra", tier: "fast", short: 2.29, long: 22.33}, + {name: "Astra ultrafast", provider: "openai", model: "gpt-6-astra", tier: "ultrafast", short: 6.87, long: 66.99}, + {name: "actual default after fast request", provider: "openai", model: "gpt-6-astra", tier: "default", cfg: latest.ModelConfig{ProviderOpts: map[string]any{"service_tier": "fast"}}, short: 1.145, long: 11.165}, + {name: "actual fast without requested tier", provider: "openai", model: "gpt-6-astra", tier: "fast", short: 2.29, long: 22.33}, + {name: "missing actual tier", provider: "openai", model: "gpt-6-astra", cfg: latest.ModelConfig{ProviderOpts: map[string]any{"service_tier": "ultrafast"}}, short: 1.145, long: 11.165}, + {name: "unknown tier", provider: "openai", model: "gpt-6-astra", tier: "future-tier", short: 1.145, long: 11.165}, + {name: "unsupported Luna ultrafast", provider: "openai", model: "gpt-5.6-luna", tier: "ultrafast", short: 0.0239, long: 0.2248}, + {name: "gateway fast already priced", provider: "vercel", model: "openai/gpt-6-astra-fast", tier: "fast", short: 2.29, long: 22.33}, + {name: "Azure unchanged", provider: "azure", model: "gpt-6-astra", tier: "fast", short: 1.145, long: 11.165}, + {name: "custom OpenAI endpoint unchanged", provider: "openai", model: "gpt-6-astra", tier: "ultrafast", cfg: latest.ModelConfig{BaseURL: "https://example.com/v1"}, short: 1.145, long: 11.165}, + {name: "override wins over ultrafast and context bands", provider: "openai", model: "gpt-6-astra", tier: "ultrafast", cfg: latest.ModelConfig{Cost: &latest.CostConfig{Input: 1, Output: 2, CacheRead: 0.1, CacheWrite: 1.25}}, short: 0.0995, long: 0.5495}, + {name: "zero override stays free", provider: "openai", model: "gpt-6-astra", tier: "fast", cfg: latest.ModelConfig{Cost: &latest.CostConfig{}}, short: 0, long: 0}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + id := modelsdev.NewID(tc.provider, tc.model) + model, err := store.GetModel(t.Context(), id) + require.NoError(t, err) + original := *model.Cost + original.Tiers = slices.Clone(original.Tiers) + for _, band := range []struct { + input int64 + want float64 + }{{50_000, tc.short}, {500_000, tc.long}} { + usage := &chat.Usage{InputTokens: band.input, CachedInputTokens: 20_000, CacheWriteTokens: 30_000, OutputTokens: 5_000, ReasoningTokens: 3_000, ServiceTier: tc.tier} + priced := applyModelCost(model, id, usage, base.Config{ModelConfig: tc.cfg}) + got := computeMessageCost(usage, priced) + require.NotNil(t, got) + assert.InDelta(t, band.want, *got, 1e-9) + assert.Equal(t, original, *model.Cost, "shared catalogue must not be mutated") + } + }) + } +} + +func TestApplyModelCostServiceTierUnpriced(t *testing.T) { + t.Parallel() + + id := modelsdev.NewID("openai", "gpt-6-astra") + usage := &chat.Usage{InputTokens: 1_000_000, ServiceTier: "ultrafast"} + assert.Nil(t, applyModelCost(nil, id, usage, base.Config{})) + model := &modelsdev.Model{} + assert.Same(t, model, applyModelCost(model, id, usage, base.Config{})) + assert.Same(t, model, applyModelCost(model, id, nil, base.Config{})) + priced := applyModelCost(nil, id, usage, base.Config{ModelConfig: latest.ModelConfig{Cost: &latest.CostConfig{Input: 1.25}}}) + cost := computeMessageCost(usage, priced) + require.NotNil(t, cost) + assert.InDelta(t, 1.25, *cost, 1e-9) +} + +func TestApplyModelCostEndpointEligibility(t *testing.T) { + t.Parallel() + + id := modelsdev.NewID("openai", "gpt-6-astra") + usage := &chat.Usage{InputTokens: 1_000_000, ServiceTier: "ultrafast"} + model := &modelsdev.Model{Cost: &modelsdev.Cost{Input: 10}} + for _, tc := range []struct { + name string + config base.Config + want float64 + }{ + {name: "official resolved endpoint", config: base.Config{BaseURL: "https://api.openai.com/v1/"}, want: 60}, + {name: "environment redirected endpoint", config: base.Config{BaseURL: "https://example.com/v1"}, want: 10}, + {name: "router endpoint unknown", config: base.Config{ModelConfig: latest.ModelConfig{Routing: []latest.RoutingRule{{Model: "custom"}}}}, want: 10}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + cost := computeMessageCost(usage, applyModelCost(model, id, usage, tc.config)) + require.NotNil(t, cost) + assert.InDelta(t, tc.want, *cost, 1e-9) + }) + } +} diff --git a/pkg/runtime/loop.go b/pkg/runtime/loop.go index 62424d8f7..79d10c8f4 100644 --- a/pkg/runtime/loop.go +++ b/pkg/runtime/loop.go @@ -24,6 +24,7 @@ import ( "github.com/docker/docker-agent/pkg/config/latest" "github.com/docker/docker-agent/pkg/httpclient" "github.com/docker/docker-agent/pkg/model/provider" + "github.com/docker/docker-agent/pkg/model/provider/base" "github.com/docker/docker-agent/pkg/modelsdev" "github.com/docker/docker-agent/pkg/runtime/toolexec" "github.com/docker/docker-agent/pkg/session" @@ -949,8 +950,8 @@ func (r *LocalRuntime) runTurn( } events.Emit(AgentInfo(a.Name(), modelID.String(), a.Description(), a.WelcomeMessage())) } - // Fallbacks may share an ID but have different endpoint pricing overrides. - m = applyConfigCost(m, modelID, usedModel.BaseConfig().ModelConfig.Cost) + // Fallbacks may share an ID but have different endpoints or cost overrides. + m = applyModelCost(m, modelID, res.Usage, usedModel.BaseConfig()) } // A successful model call resets the overflow compaction counter. @@ -1265,6 +1266,22 @@ func (r *LocalRuntime) Run(ctx context.Context, sess *session.Session) ([]sessio return sess.GetAllMessages(), nil } +func applyModelCost(m *modelsdev.Model, id modelsdev.ID, usage *chat.Usage, config base.Config) *modelsdev.Model { + cfg := config.ModelConfig + // Routers cannot identify the serving endpoint here; custom endpoints need their own rates. + customEndpoint := cfg.BaseURL != "" || (config.BaseURL != "" && strings.TrimRight(config.BaseURL, "/") != "https://api.openai.com/v1") + if cfg.Cost != nil || customEndpoint || len(cfg.Routing) > 0 || m == nil || usage == nil { + return applyConfigCost(m, id, cfg.Cost) + } + cost := m.Cost.ForServiceTier(id, usage.ServiceTier) + if cost == m.Cost { + return m + } + out := *m + out.Cost = cost + return &out +} + // applyConfigCost overlays a config-declared price table (USD per 1M tokens) // onto the catalogue entry, returning m untouched when there is no override. // It never mutates m: the store caches entries shared across sessions. When diff --git a/pkg/runtime/native_compaction.go b/pkg/runtime/native_compaction.go index 3fbb25f22..a14b2df0e 100644 --- a/pkg/runtime/native_compaction.go +++ b/pkg/runtime/native_compaction.go @@ -123,7 +123,7 @@ func (r *LocalRuntime) compactNatively(ctx context.Context, sess *session.Sessio slog.DebugContext(ctx, "Failed to get model definition for native compaction cost", "model_id", modelID.String(), "error", err) m = nil } - m = applyConfigCost(m, modelID, native.BaseConfig().ModelConfig.Cost) + m = applyModelCost(m, modelID, &res.Usage, native.BaseConfig()) messageCost := computeMessageCost(&res.Usage, m) r.recordBudget(sess, a, &res.Usage, messageCost, r.now().Sub(started), events) if strings.TrimSpace(res.Summary) == "" { diff --git a/pkg/runtime/service_tier_cost_test.go b/pkg/runtime/service_tier_cost_test.go new file mode 100644 index 000000000..9d63bc612 --- /dev/null +++ b/pkg/runtime/service_tier_cost_test.go @@ -0,0 +1,104 @@ +package runtime + +import ( + "context" + "errors" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/agent" + "github.com/docker/docker-agent/pkg/chat" + "github.com/docker/docker-agent/pkg/config/latest" + "github.com/docker/docker-agent/pkg/hooks" + "github.com/docker/docker-agent/pkg/modelsdev" + "github.com/docker/docker-agent/pkg/session" + "github.com/docker/docker-agent/pkg/team" +) + +func TestRunStreamServiceTierCost(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name, model, tier string + fallback bool + cost *latest.CostConfig + want float64 + }{ + {name: "Luna priority", model: "gpt-5.6-luna", tier: "priority", want: 0.0478}, + {name: "Astra fast", model: "gpt-6-astra", tier: "fast", want: 2.29}, + {name: "Astra ultrafast", model: "gpt-6-astra", tier: "ultrafast", want: 6.87}, + {name: "standard downgrade", model: "gpt-6-astra", tier: "default", want: 1.145}, + {name: "fallback ultrafast", model: "gpt-6-astra", tier: "ultrafast", fallback: true, want: 6.87}, + {name: "fallback override", model: "gpt-6-astra", tier: "ultrafast", fallback: true, cost: &latest.CostConfig{Input: 1, Output: 2, CacheRead: 0.1, CacheWrite: 1.25}, want: 0.0995}, + {name: "free override", model: "gpt-6-astra", tier: "fast", cost: &latest.CostConfig{}, want: 0}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + usage := &chat.Usage{InputTokens: 50_000, CachedInputTokens: 20_000, CacheWriteTokens: 30_000, OutputTokens: 5_000, ReasoningTokens: 3_000, ServiceTier: tc.tier} + stream := newStreamBuilder().AddContent("ok").AddStopWithUsage(0, 0).Build() + stream.responses[len(stream.responses)-1].Usage = usage + id := "openai/" + tc.model + model := &pricingProvider{Provider: &mockProvider{id: id, stream: stream}, cost: tc.cost} + opts := []agent.Opt{ + agent.WithModel(model), + agent.WithHooks(&latest.HooksConfig{AfterLLMCall: []latest.HookDefinition{{Type: "builtin", Command: "capture-tier-cost"}}}), + } + if tc.fallback { + opts = append(opts, + agent.WithModel(&countingProvider{id: "openai/gpt-5.6-luna", failCount: 100, err: errors.New("401 unauthorized")}), + agent.WithFallbackModel(model), agent.WithFallbackRetries(-1)) + } + root := agent.New("root", "test", opts...) + rec := &recordingTelemetry{} + rt, err := NewLocalRuntime(t.Context(), team.New(team.WithAgents(root)), + WithModelStore(modelsdev.NewDatabaseStore(modelsdev.EmbeddedSnapshot())), + WithSessionCompaction(false), WithTelemetry(rec), WithBudget(&latest.BudgetConfig{MaxCost: 100})) + require.NoError(t, err) + var captured *hooks.Input + require.NoError(t, rt.hooksRegistry.RegisterBuiltin("capture-tier-cost", + func(_ context.Context, in *hooks.Input, _ []string) (*hooks.Output, error) { + snapshot := *in + captured = &snapshot + return nil, nil + })) + sess := session.New(session.WithUserMessage("hi")) + sess.Title = "Service tier test" + var lastUsage *MessageUsage + var budget *BudgetStatus + for ev := range rt.RunStream(t.Context(), sess) { + switch ev := ev.(type) { + case *ErrorEvent: + t.Errorf("unexpected runtime error: %s", ev.Error) + case *TokenUsageEvent: + if ev.Usage != nil && ev.Usage.LastMessage != nil { + lastUsage = ev.Usage.LastMessage + } + case *BudgetUsageEvent: + for _, b := range ev.Budgets { + if b.Name == runBudgetName { + budget = &b + } + } + } + } + require.NotNil(t, captured) + require.NotNil(t, captured.Cost) + assert.Equal(t, id, captured.ModelID) + require.NotNil(t, captured.Usage) + assert.Equal(t, tc.tier, captured.Usage.ServiceTier) + assert.InDelta(t, tc.want, *captured.Cost, 1e-9) + assert.InDelta(t, tc.want, sess.OwnCost(), 1e-9) + require.NotNil(t, lastUsage) + assert.InDelta(t, tc.want, lastUsage.Cost, 1e-9) + require.NotNil(t, budget) + assert.InDelta(t, tc.want, budget.Cost, 1e-9) + records := rec.snapshot().tokenUsages + require.Len(t, records, 1) + assert.InDelta(t, tc.want, records[0].Cost, 1e-9) + assert.Equal(t, int64(100_000), records[0].InputTokens) + assert.Equal(t, int64(5_000), records[0].OutputTokens) + }) + } +}