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) + }) + } +}