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
8 changes: 6 additions & 2 deletions docs/providers/openai/index.md
Original file line number Diff line number Diff line change
Expand Up @@ -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):

Expand All @@ -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.

Expand Down
11 changes: 8 additions & 3 deletions examples/openai-service-tier.yaml
Original file line number Diff line number Diff line change
@@ -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.

Expand All @@ -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
4 changes: 3 additions & 1 deletion pkg/chat/chat.go
Original file line number Diff line number Diff line change
Expand Up @@ -216,14 +216,16 @@ 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.
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 {
Expand Down
18 changes: 18 additions & 0 deletions pkg/chat/chat_test.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package chat

import (
"encoding/json"
"fmt"
"os"
"path/filepath"
Expand Down Expand Up @@ -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")
}
5 changes: 5 additions & 0 deletions pkg/model/provider/oaistream/adapter.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand All @@ -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{
Expand Down Expand Up @@ -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
Expand Down
49 changes: 49 additions & 0 deletions pkg/model/provider/oaistream/adapter_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
})
}
}
9 changes: 7 additions & 2 deletions pkg/model/provider/openai/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ import (
"fmt"
"log/slog"
"net/http"
"os"
"slices"
"strings"
"sync"
Expand Down Expand Up @@ -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 == "" {
Expand Down Expand Up @@ -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,
}
Expand All @@ -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
Expand Down
6 changes: 6 additions & 0 deletions pkg/model/provider/openai/response_stream.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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{}

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
58 changes: 58 additions & 0 deletions pkg/model/provider/openai/response_stream_terminal_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
})
}
}
}
Loading
Loading