diff --git a/docs/community/troubleshooting/index.md b/docs/community/troubleshooting/index.md index f81f50cc94..650f3feed5 100644 --- a/docs/community/troubleshooting/index.md +++ b/docs/community/troubleshooting/index.md @@ -287,6 +287,12 @@ A few things that catch people out: > [!WARNING] > Raising `--max-request-size` increases how much memory an unauthenticated or malicious client can force the server to buffer per request. Pick a value with your deployment's exposure in mind, and pair any non-loopback listener with `--auth-token` (API server) or `--api-key`/`--api-key-env` (chat server). A `--listen` control plane has neither flag — keep it on loopback, a unix socket, or behind an authenticating reverse proxy if it must be reachable from elsewhere. +### Delegated task appears stalled before its first response + +A direct `transfer_task` delegation automatically retries once when the child model stream is silent before sending any response payload. It emits the existing warning event, and the retry is immediate and applies once to the complete child run, including later turns and fallback models. Each nested native `transfer_task` starts a new child run with its own fresh retry allowance. The retry does not apply after partial output, to `background_agents`, or when cancellation or a run budget blocks it. + +Run with `--debug` and look for `Delegated model stream idle before response; retrying immediately`. If the child still fails, check provider connectivity, fallback configuration, cancellation, and run-budget events; Docker Agent does not retry when the context is canceled or the run budget is exhausted. + ## Performance Issues ### High memory usage diff --git a/docs/configuration/models/index.md b/docs/configuration/models/index.md index 7e7656c1e8..e887e5f7c9 100644 --- a/docs/configuration/models/index.md +++ b/docs/configuration/models/index.md @@ -32,7 +32,7 @@ models: token_key: string # Optional: env var for API token thinking_budget: string|int # Optional: reasoning effort task_budget: int|object # Optional: total task token budget (Anthropic) - parallel_tool_calls: boolean # Optional: allow parallel tool calls + parallel_tool_calls: boolean # Optional: allow parallel tool calls. Omit to use the provider/API default. track_usage: boolean # Optional: track token usage routing: [list] # Optional: rule-based model routing capabilities: # Optional: override attachment (input) capabilities @@ -72,7 +72,7 @@ models: | `token_key` | string | ✗ | Environment variable name containing the API token (overrides provider default) | | `thinking_budget` | string/int | ✗ | Reasoning effort control | | `task_budget` | int/object | ✗ | Total token budget for an agentic task (forwarded to Anthropic; see [Task Budget](#task-budget)). | -| `parallel_tool_calls` | boolean | ✗ | Allow model to call multiple tools at once | +| `parallel_tool_calls` | boolean | ✗ | Allow model to call multiple tools at once. When omitted, Docker Agent leaves the setting unset so the selected provider or API can apply its own default. | | `track_usage` | boolean | ✗ | Track and report token usage for this model | | `routing` | array | ✗ | Rule-based routing to different models. See [Model Routing](../routing/index.md). | | `capabilities` | object | ✗ | Override attachment (input) capabilities for this model. See [Attachment Capability Overrides](#attachment-capability-overrides). | diff --git a/docs/configuration/overview/index.md b/docs/configuration/overview/index.md index 0f5f33af85..5c65985c37 100644 --- a/docs/configuration/overview/index.md +++ b/docs/configuration/overview/index.md @@ -485,7 +485,7 @@ agents: | `top_p` | Default top-p sampling parameter. | | `frequency_penalty` | Default frequency penalty. | | `presence_penalty` | Default presence penalty. | -| `parallel_tool_calls` | Enable parallel tool calls by default. | +| `parallel_tool_calls` | Enable or disable parallel tool calls by default. If omitted, the provider/API default is used. | | `track_usage` | Track token usage by default. | | `provider_opts` | Provider-specific options. | diff --git a/docs/features/cli/index.md b/docs/features/cli/index.md index c57db6075b..d7f27d77bf 100644 --- a/docs/features/cli/index.md +++ b/docs/features/cli/index.md @@ -30,7 +30,7 @@ $ docker agent run [config] [message...] [flags] | `-a, --agent ` | Run a specific agent from the config | | `--yolo` | Auto-approve tool calls (unless explicitly denied). Legacy alias for `--safety autonomous`. | | `--safety ` | Safety mode for tool approval: `strict` (ask for everything), `balanced` (auto-approve safe calls), `restricted` (auto-approve safe calls, deny the rest — fail-closed for unattended runs), or `autonomous` (approve everything). Wins over `--yolo` when both are given. Without the flag, the mode falls back to alias/user-config defaults, then the agent YAML's `agents..safety` / `runtime.safety`; a resumed session keeps its stored mode unless `--safety`/`--yolo` is passed explicitly. See [Safety Modes](../../configuration/permissions/index.md#safety-modes). | -| `--model ` | Override model(s). Use `provider/model` for all agents, or `agent=provider/model` for specific agents. Comma-separate multiple overrides. | +| `--model ` | Override model(s). Use `provider/model` for all agents, or `agent=provider/model` for specific agents. Comma-separate multiple overrides. Inline overrides inherit each agent's explicit `parallel_tool_calls` setting only from a concrete, provider-qualified configured model. Existing named targets are authoritative on every invocation, including an omitted setting; harness, router, alloy/comma-list, `first_available`, providerless, and unresolved source models do not inherit it. | | `--session ` | Resume a previous session. Supports relative refs (`-1` = newest by creation time, `-2` = second-newest, … — creation order, not last-used). An explicit ID that does not exist yet is created with that ID, so a supervisor can own the session ID upfront and reuse it across runs. | | `-s, --session-db ` | Path to the SQLite session database (default: `/session.db`, so `~/.cagent/session.db` unless `--data-dir` is set) | | `--session-read-only` | Open the TUI in read-only mode: conversation history is displayed but no new messages can be sent to the LLM. Cannot be used with `--exec`. | diff --git a/docs/providers/custom/index.md b/docs/providers/custom/index.md index d664ca3d84..4cf97522d5 100644 --- a/docs/providers/custom/index.md +++ b/docs/providers/custom/index.md @@ -104,7 +104,7 @@ agents: | `top_p` | float | Default nucleus sampling threshold (0.0–1.0). | — | | `frequency_penalty` | float | Default frequency penalty (-2.0–2.0). | — | | `presence_penalty` | float | Default presence penalty (-2.0–2.0). | — | -| `parallel_tool_calls` | boolean | Whether to enable parallel tool calls by default. | — | +| `parallel_tool_calls` | boolean | Whether to enable parallel tool calls by default. When omitted, the provider/API default is used. | — | | `track_usage` | boolean | Whether to track token usage by default. | — | | `thinking_budget` | string/int | Default reasoning effort/budget. | — | | `task_budget` | int/object | Default total token budget for an agentic task (forwarded to Anthropic; honored by Claude Opus 4.7+ today). Integer shorthand or `{type: tokens, total: N}`. | — | diff --git a/docs/tools/transfer-task/index.md b/docs/tools/transfer-task/index.md index 9d3bc2a03a..ef4831bb51 100644 --- a/docs/tools/transfer-task/index.md +++ b/docs/tools/transfer-task/index.md @@ -15,6 +15,10 @@ The `transfer_task` tool allows an agent to delegate tasks to specialized sub-ag **You don't need to add it manually** — it's automatically available when an agent has `sub_agents` configured. +## Idle stream recovery + +When a delegated model connection becomes idle before producing any response payload, `transfer_task` retries that stream once and emits a warning. The retry allowance belongs to the complete direct delegated child run: it can be consumed only once across all of that child's turns and fallback models. A nested native `transfer_task` starts its own child run with a fresh allowance rather than inheriting the caller's; background delegation has no allowance. The retry is skipped when cancellation or a run budget blocks it and is never attempted after partial content, reasoning, media, or tool-call data has arrived, preventing duplicate output or tool execution. + ## Configuration The tool is enabled implicitly when `sub_agents` is set: diff --git a/pkg/config/config.go b/pkg/config/config.go index f85a851011..cd5f2077c8 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -262,14 +262,6 @@ func validateConfig(cfg *latest.Config) error { cfg.Models = map[string]latest.ModelConfig{} } - for name := range cfg.Models { - if cfg.Models[name].ParallelToolCalls == nil { - m := cfg.Models[name] - m.ParallelToolCalls = new(true) - cfg.Models[name] = m - } - } - if err := ensureModelsExist(cfg); err != nil { return err } diff --git a/pkg/config/latest/internal_state.go b/pkg/config/latest/internal_state.go new file mode 100644 index 0000000000..b54d5a9dbc --- /dev/null +++ b/pkg/config/latest/internal_state.go @@ -0,0 +1,18 @@ +package latest + +// InternalModelOverrideState returns opaque loader-owned state. It is not part +// of the configuration schema and must not affect serialization. +func (t *Config) InternalModelOverrideState() any { + if t == nil { + return nil + } + return t.modelOverrideState +} + +// SetInternalModelOverrideState stores opaque loader-owned state outside the +// public configuration schema. +func (t *Config) SetInternalModelOverrideState(state any) { + if t != nil { + t.modelOverrideState = state + } +} diff --git a/pkg/config/latest/types.go b/pkg/config/latest/types.go index 7c73a8668b..43a9f9ac25 100644 --- a/pkg/config/latest/types.go +++ b/pkg/config/latest/types.go @@ -65,6 +65,8 @@ type Config struct { // only carries the section so it round-trips — see config.applyFlavors // for the merge semantics. Flavors map[string]map[string]any `json:"flavors,omitempty"` + + modelOverrideState any } // BudgetConfig caps what a single run may consume before the agent is diff --git a/pkg/config/overrides.go b/pkg/config/overrides.go index 735e124d8a..cc856bbd0a 100644 --- a/pkg/config/overrides.go +++ b/pkg/config/overrides.go @@ -3,31 +3,163 @@ package config import ( "errors" "fmt" + "reflect" + "slices" "strings" "github.com/docker/docker-agent/pkg/config/latest" ) -// ApplyModelOverrides applies CLI model overrides to the configuration +type modelOverrideReceipt struct { + sourceModel string + sourcePolicy *bool + eligible bool + requestedModel string + expectedPostModel string + expectedPostValue latest.ModelConfig +} + +type modelOverridePolicy struct { + modelRef string + parallelToolCalls *bool +} + +type modelOverrideState struct { + receipts map[string]modelOverrideReceipt + policies map[string]modelOverridePolicy + lastOverrides []string +} + +func overrideState(cfg *latest.Config) *modelOverrideState { + state, _ := cfg.InternalModelOverrideState().(*modelOverrideState) + if state == nil { + state = &modelOverrideState{ + receipts: make(map[string]modelOverrideReceipt), + policies: make(map[string]modelOverridePolicy), + } + cfg.SetInternalModelOverrideState(state) + } + return state +} + +func cloneBool(value *bool) *bool { + if value == nil { + return nil + } + return new(*value) +} + +func isConcreteModelConfig(model latest.ModelConfig) bool { + return model.Provider != "" && model.Model != "" && !strings.Contains(model.Model, ",") && + len(model.Routing) == 0 && !model.IsFirstAvailable() +} + +func sourceReceipt(cfg *latest.Config, state *modelOverrideState, agent latest.AgentConfig) modelOverrideReceipt { + if receipt, ok := state.receipts[agent.Name]; ok { + current, exists := cfg.Models[agent.Model] + if agent.Model == receipt.expectedPostModel && exists && reflect.DeepEqual(current, receipt.expectedPostValue) { + return receipt + } + delete(state.receipts, agent.Name) + } + + model, exists := cfg.Models[agent.Model] + return modelOverrideReceipt{ + sourceModel: agent.Model, + sourcePolicy: cloneBool(model.ParallelToolCalls), + eligible: agent.Harness == nil && exists && isConcreteModelConfig(model), + } +} + +// ApplyModelOverridePolicy applies an inherited per-agent CLI override policy +// to a copied model config. The manifest model map remains authoritative and +// unmodified. +func ApplyModelOverridePolicy(cfg *latest.Config, agentName, modelRef string, model *latest.ModelConfig) { + if cfg == nil || model == nil { + return + } + state, _ := cfg.InternalModelOverrideState().(*modelOverrideState) + if state == nil { + return + } + policy, ok := state.policies[agentName] + if !ok || policy.modelRef != modelRef { + return + } + model.ParallelToolCalls = cloneBool(policy.parallelToolCalls) +} + +// ApplyModelOverrides applies CLI model overrides to the configuration. func ApplyModelOverrides(cfg *latest.Config, overrides []string) error { + if len(overrides) == 0 { + return nil + } + + state := overrideState(cfg) + if slices.Equal(overrides, state.lastOverrides) && receiptsIntact(cfg, state) { + return nil + } + invocationTargets := make(map[string]struct{}, len(cfg.Models)) + for name := range cfg.Models { + invocationTargets[name] = struct{}{} + } + receipts := make(map[string]modelOverrideReceipt, len(cfg.Agents)) + for _, agent := range cfg.Agents { + receipts[agent.Name] = sourceReceipt(cfg, state, agent) + } + for _, override := range overrides { if err := applySingleOverride(cfg, override); err != nil { return err } } - // After applying overrides, ensure new models are added to cfg.Models - return ensureModelsExist(cfg) + if err := ensureModelsExist(cfg); err != nil { + return err + } + policies := make(map[string]modelOverridePolicy) + for _, agent := range cfg.Agents { + receipt := receipts[agent.Name] + requestedRef := agent.Model + target := cfg.Models[requestedRef] + _, targetExisted := invocationTargets[requestedRef] + if receipt.eligible && receipt.sourcePolicy != nil && !targetExisted { + policies[agent.Name] = modelOverridePolicy{ + modelRef: requestedRef, + parallelToolCalls: cloneBool(receipt.sourcePolicy), + } + } + receipt.requestedModel = requestedRef + receipt.expectedPostModel = requestedRef + receipt.expectedPostValue = target + state.receipts[agent.Name] = receipt + } + state.policies = policies + state.lastOverrides = slices.Clone(overrides) + return nil +} + +func receiptsIntact(cfg *latest.Config, state *modelOverrideState) bool { + for _, agent := range cfg.Agents { + receipt, ok := state.receipts[agent.Name] + if !ok || agent.Model != receipt.expectedPostModel { + return false + } + model, exists := cfg.Models[agent.Model] + if !exists || !reflect.DeepEqual(model, receipt.expectedPostValue) { + return false + } + } + return len(state.receipts) == len(cfg.Agents) } -// applySingleOverride processes a single model override string +// applySingleOverride processes a single model override string. func applySingleOverride(cfg *latest.Config, override string) error { override = strings.TrimSpace(override) if override == "" { - return nil // Skip empty overrides + return nil } - // Handle comma-separated format: "agent1=model1,agent2=model2" if strings.Contains(override, ",") { for part := range strings.SplitSeq(override, ",") { if err := applySingleOverride(cfg, part); err != nil { @@ -37,7 +169,6 @@ func applySingleOverride(cfg *latest.Config, override string) error { return nil } - // Check if this is an agent-specific override (contains '=') agentName, modelSpec, ok := strings.Cut(override, "=") if ok { agentName = strings.TrimSpace(agentName) @@ -50,7 +181,6 @@ func applySingleOverride(cfg *latest.Config, override string) error { return fmt.Errorf("empty model specification in override: %s", override) } - // Apply to specific agent ok := cfg.Agents.Update(agentName, func(a *latest.AgentConfig) { a.Model = modelSpec }) @@ -58,7 +188,6 @@ func applySingleOverride(cfg *latest.Config, override string) error { return fmt.Errorf("unknown agent '%s'", agentName) } } else { - // Global override: apply to all agents modelSpec := strings.TrimSpace(override) if modelSpec == "" { return errors.New("empty model specification") diff --git a/pkg/config/overrides_test.go b/pkg/config/overrides_test.go new file mode 100644 index 0000000000..3871319d63 --- /dev/null +++ b/pkg/config/overrides_test.go @@ -0,0 +1,246 @@ +package config + +import ( + "maps" + "slices" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/config/latest" +) + +func TestApplyModelOverrides_InheritsParallelToolCalls(t *testing.T) { + boolPtr := func(value bool) *bool { return &value } + tests := []struct { + name string + agents []latest.AgentConfig + models map[string]latest.ModelConfig + overrides []string + wantModels map[string]string + wantValues map[string]*bool + }{ + { + name: "inherits false into inline target", + agents: []latest.AgentConfig{{Name: "root", Model: "configured"}}, + models: map[string]latest.ModelConfig{"configured": {Provider: "openai", Model: "configured", ParallelToolCalls: boolPtr(false)}}, + overrides: []string{"openai/replacement"}, + wantModels: map[string]string{"root": "openai/replacement"}, + wantValues: map[string]*bool{"root": boolPtr(false)}, + }, + { + name: "inherits true into inline target", + agents: []latest.AgentConfig{{Name: "root", Model: "configured"}}, + models: map[string]latest.ModelConfig{"configured": {Provider: "openai", Model: "configured", ParallelToolCalls: boolPtr(true)}}, + overrides: []string{"openai/replacement"}, + wantModels: map[string]string{"root": "openai/replacement"}, + wantValues: map[string]*bool{"root": boolPtr(true)}, + }, + { + name: "named target explicit value wins", + agents: []latest.AgentConfig{{Name: "root", Model: "configured"}}, + models: map[string]latest.ModelConfig{"configured": {Provider: "openai", Model: "configured", ParallelToolCalls: boolPtr(false)}, "replacement": {Provider: "openai", Model: "replacement", ParallelToolCalls: boolPtr(true)}}, + overrides: []string{"replacement"}, + wantModels: map[string]string{"root": "replacement"}, + wantValues: map[string]*bool{"root": boolPtr(true)}, + }, + { + name: "named target omitted value wins", + agents: []latest.AgentConfig{{Name: "root", Model: "configured"}}, + models: map[string]latest.ModelConfig{"configured": {Provider: "openai", Model: "configured", ParallelToolCalls: boolPtr(false)}, "replacement": {Provider: "openai", Model: "replacement"}}, + overrides: []string{"replacement"}, + wantModels: map[string]string{"root": "replacement"}, + wantValues: map[string]*bool{"root": nil}, + }, + { + name: "source omitted does not inherit", + agents: []latest.AgentConfig{{Name: "root", Model: "configured"}}, + models: map[string]latest.ModelConfig{"configured": {Provider: "openai", Model: "configured"}}, + overrides: []string{"openai/replacement"}, + wantModels: map[string]string{"root": "openai/replacement"}, + wantValues: map[string]*bool{"root": nil}, + }, + { + name: "shared target isolates per-agent policy", + agents: []latest.AgentConfig{ + {Name: "root", Model: "serial"}, {Name: "worker", Model: "parallel"}, {Name: "unset", Model: "unset"}, + }, + models: map[string]latest.ModelConfig{ + "serial": {Provider: "openai", Model: "serial", ParallelToolCalls: boolPtr(false)}, + "parallel": {Provider: "openai", Model: "parallel", ParallelToolCalls: boolPtr(true)}, + "unset": {Provider: "openai", Model: "unset"}, + }, + overrides: []string{"openai/replacement"}, + wantModels: map[string]string{"root": "openai/replacement", "worker": "openai/replacement", "unset": "openai/replacement"}, + wantValues: map[string]*bool{"root": boolPtr(false), "worker": boolPtr(true), "unset": nil}, + }, + { + name: "excludes harness router alloy selector unresolved and providerless sources", + agents: []latest.AgentConfig{ + {Name: "harness", Harness: &latest.HarnessConfig{Type: "codex"}}, + {Name: "router", Model: "router"}, + {Name: "alloy", Model: "alloy"}, + {Name: "selector", Model: "selector"}, + {Name: "missing", Model: "missing"}, + {Name: "providerless", Model: "providerless"}, + {Name: "csv", Model: "csv"}, + }, + models: map[string]latest.ModelConfig{ + "router": {Provider: "openai", Routing: []latest.RoutingRule{{Model: "leaf"}}, ParallelToolCalls: boolPtr(false)}, + "alloy": {Model: "leaf,leaf2", ParallelToolCalls: boolPtr(false)}, + "selector": {Provider: "openai", FirstAvailable: []string{"leaf"}, ParallelToolCalls: boolPtr(false)}, + "providerless": {Model: "leaf", ParallelToolCalls: boolPtr(false)}, + "csv": {Provider: "openai", Model: "leaf,leaf2", ParallelToolCalls: boolPtr(false)}, + "leaf": {Provider: "openai", Model: "leaf"}, "leaf2": {Provider: "openai", Model: "leaf2"}, + }, + overrides: []string{"openai/replacement"}, + wantModels: map[string]string{ + "harness": "openai/replacement", "router": "openai/replacement", "alloy": "openai/replacement", + "selector": "openai/replacement", "missing": "openai/replacement", "providerless": "openai/replacement", "csv": "openai/replacement", + }, + wantValues: map[string]*bool{"harness": nil, "router": nil, "alloy": nil, "selector": nil, "missing": nil, "providerless": nil, "csv": nil}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cfg := &latest.Config{Agents: slices.Clone(tt.agents), Models: maps.Clone(tt.models)} + require.NoError(t, ApplyModelOverrides(cfg, tt.overrides)) + expectedModelCount := len(tt.models) + for _, modelRef := range tt.wantModels { + if _, existed := tt.models[modelRef]; !existed { + expectedModelCount++ + break + } + } + assert.Len(t, cfg.Models, expectedModelCount) + for name, original := range tt.models { + assert.Equal(t, original, cfg.Models[name], "pre-existing model %q mutated", name) + } + for _, agent := range cfg.Agents { + assert.Equal(t, tt.wantModels[agent.Name], agent.Model) + model := cfg.Models[agent.Model] + ApplyModelOverridePolicy(cfg, agent.Name, agent.Model, &model) + want := tt.wantValues[agent.Name] + if want == nil { + assert.Nil(t, model.ParallelToolCalls) + } else if assert.NotNil(t, model.ParallelToolCalls) { + assert.Equal(t, *want, *model.ParallelToolCalls) + } + } + for name := range cfg.Models { + assert.NotContains(t, name, "__cli_model_") + } + }) + } +} + +func TestApplyModelOverrides_RepeatedSameOverrideIsIdempotent(t *testing.T) { + falseValue := false + cfg := &latest.Config{ + Agents: []latest.AgentConfig{{Name: "root", Model: "configured"}}, + Models: map[string]latest.ModelConfig{ + "configured": {Provider: "openai", Model: "configured", ParallelToolCalls: &falseValue}, + }, + } + + require.NoError(t, ApplyModelOverrides(cfg, []string{"openai/first"})) + firstState := cfg.InternalModelOverrideState() + modelCount := len(cfg.Models) + + require.NoError(t, ApplyModelOverrides(cfg, []string{"openai/first"})) + assert.Same(t, firstState, cfg.InternalModelOverrideState()) + assert.Equal(t, "openai/first", cfg.Agents[0].Model) + assert.Len(t, cfg.Models, modelCount) + model := cfg.Models["openai/first"] + ApplyModelOverridePolicy(cfg, "root", "openai/first", &model) + require.NotNil(t, model.ParallelToolCalls) + assert.False(t, *model.ParallelToolCalls) +} + +func TestApplyModelOverrides_FreshInvocationAuthorityAndDrift(t *testing.T) { + falseValue := false + trueValue := true + cfg := &latest.Config{ + Agents: []latest.AgentConfig{{Name: "root", Model: "configured"}}, + Models: map[string]latest.ModelConfig{ + "configured": {Provider: "openai", Model: "configured", ParallelToolCalls: &falseValue}, + }, + } + + require.NoError(t, ApplyModelOverrides(cfg, []string{"openai/first"})) + assert.Equal(t, "openai/first", cfg.Agents[0].Model) + + // A raw target created by the previous call is authoritative on a changed call. + require.NoError(t, ApplyModelOverrides(cfg, []string{"root=openai/first"})) + model := cfg.Models["openai/first"] + ApplyModelOverridePolicy(cfg, "root", "openai/first", &model) + assert.Nil(t, model.ParallelToolCalls) + + // A target added after the first call is authoritative on the next call. + cfg.Models["named"] = latest.ModelConfig{Provider: "openai", Model: "named"} + require.NoError(t, ApplyModelOverrides(cfg, []string{"named"})) + model = cfg.Models["named"] + ApplyModelOverridePolicy(cfg, "root", "named", &model) + assert.Nil(t, model.ParallelToolCalls) + + // Unchanged expected post-state keeps the original source receipt. + require.NoError(t, ApplyModelOverrides(cfg, []string{"openai/second"})) + model = cfg.Models["openai/second"] + ApplyModelOverridePolicy(cfg, "root", "openai/second", &model) + require.NotNil(t, model.ParallelToolCalls) + assert.False(t, *model.ParallelToolCalls) + + // External drift invalidates the receipt and snapshots the current source. + drifted := cfg.Models["openai/second"] + drifted.ParallelToolCalls = &trueValue + cfg.Models["openai/second"] = drifted + require.NoError(t, ApplyModelOverrides(cfg, []string{"openai/third"})) + model = cfg.Models["openai/third"] + ApplyModelOverridePolicy(cfg, "root", "openai/third", &model) + require.NotNil(t, model.ParallelToolCalls) + assert.True(t, *model.ParallelToolCalls) +} + +func TestApplyModelOverridePolicyRequiresAgentAndRealRef(t *testing.T) { + falseValue := false + cfg := &latest.Config{ + Agents: []latest.AgentConfig{{Name: "root", Model: "configured"}}, + Models: map[string]latest.ModelConfig{ + "configured": {Provider: "openai", Model: "configured", ParallelToolCalls: &falseValue}, + }, + } + require.NoError(t, ApplyModelOverrides(cfg, []string{"openai/replacement"})) + + for _, tc := range []struct { + agent, ref string + wantNil bool + }{ + {agent: "root", ref: "openai/replacement"}, + {agent: "other", ref: "openai/replacement", wantNil: true}, + {agent: "root", ref: "openai/other", wantNil: true}, + } { + model := latest.ModelConfig{Provider: "openai", Model: "replacement"} + ApplyModelOverridePolicy(cfg, tc.agent, tc.ref, &model) + if tc.wantNil { + assert.Nil(t, model.ParallelToolCalls) + } else if assert.NotNil(t, model.ParallelToolCalls) { + assert.False(t, *model.ParallelToolCalls) + } + } +} + +func TestIsConcreteModelConfigRequiresProviderAndModel(t *testing.T) { + t.Parallel() + + assert.True(t, isConcreteModelConfig(latest.ModelConfig{Provider: "openai", Model: "gpt-5"})) + assert.False(t, isConcreteModelConfig(latest.ModelConfig{Provider: "openai"})) + assert.False(t, isConcreteModelConfig(latest.ModelConfig{Model: "gpt-5"})) +} + +func TestApplyModelOverrides_NoOverridesPreservesNilModels(t *testing.T) { + cfg := &latest.Config{} + require.NoError(t, ApplyModelOverrides(cfg, nil)) + assert.Nil(t, cfg.Models) +} diff --git a/pkg/runtime/agent_delegation.go b/pkg/runtime/agent_delegation.go index b93e5ef6ac..1f935e6ca8 100644 --- a/pkg/runtime/agent_delegation.go +++ b/pkg/runtime/agent_delegation.go @@ -206,6 +206,9 @@ type delegationRequest struct { // concurrent foreground loop and must not be mutated from a // background task (#3886). SwitchCurrentAgent bool + // directNativeTransfer enables the idle-stream recovery reserved for the + // native transfer_task handler. Other runForwarding callers leave it false. + directNativeTransfer bool } // newSubSession builds a *session.Session from a SubSessionConfig and a parent @@ -375,18 +378,24 @@ func (r *LocalRuntime) runForwarding(ctx context.Context, parent *session.Sessio // subagent_stop fires after the child's stream has fully drained, // using the *parent* agent's executor so handlers configured on the // orchestrator see every child completion in one place — success or - // failure. The deferred call ensures we don't lose the event when an - // ErrorEvent triggers an early return below; handlers can detect a - // failed run by an empty stop_response (or by correlating with the - // session-level error event the parent already received). + // failure. On failure, stop_response carries any assistant content the + // child produced before stopping; an empty value means no content existed. defer func() { r.executeSubagentStopHooks(ctx, parent, s, callerAgent, req.AgentName, s.GetLastAssistantMessageContent()) }() - childEvents := r.RunStream(ctx, s) + idleRetryPolicy := defaultIdleStreamRetryPolicy() + if req.directNativeTransfer { + idleRetryPolicy = idleStreamRetryPolicy{enabled: true, parentSessionID: parent.ID} + } + childEvents := r.runStream(ctx, s, idleRetryPolicy) var subSessionErr error for event := range childEvents { evts.Emit(event) + if budgetEvent, ok := event.(*BudgetExceededEvent); ok && + req.directNativeTransfer && budgetEvent.SessionID == s.ID && subSessionErr == nil { + subSessionErr = errors.New(budgetEvent.Message) + } if errEvent, ok := event.(*ErrorEvent); ok && subSessionErr == nil { // Capture the first ErrorEvent but keep draining the channel so // the sub-session's full transcript still streams through. The @@ -405,6 +414,9 @@ func (r *LocalRuntime) runForwarding(ctx context.Context, parent *session.Sessio parent.AddLiveSubSession(s) evts.Emit(SubSessionCompleted(parent.ID, s, callerAgent.Name())) + if subSessionErr == nil && ctx.Err() != nil { + subSessionErr = ctx.Err() + } if subSessionErr != nil { span.RecordError(subSessionErr) span.SetStatus(codes.Error, "sub-session error") @@ -739,7 +751,8 @@ func (r *LocalRuntime) handleTaskTransfer(ctx context.Context, sess *session.Ses NonInteractive: sess.NonInteractive, DelegationLineage: childLineage, }, - SwitchCurrentAgent: true, + SwitchCurrentAgent: true, + directNativeTransfer: true, }) } diff --git a/pkg/runtime/agent_delegation_test.go b/pkg/runtime/agent_delegation_test.go index c9af233811..7dba37109b 100644 --- a/pkg/runtime/agent_delegation_test.go +++ b/pkg/runtime/agent_delegation_test.go @@ -1,21 +1,27 @@ package runtime import ( + "bytes" "context" "fmt" + "log/slog" "path/filepath" "strings" "sync" "testing" + "testing/synctest" "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel/trace/noop" "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/model/provider/base" + "github.com/docker/docker-agent/pkg/modelsdev" "github.com/docker/docker-agent/pkg/permissions" "github.com/docker/docker-agent/pkg/runtime/toolexec" "github.com/docker/docker-agent/pkg/safety" @@ -637,6 +643,406 @@ func TestRunAgent_EndToEndPermissions(t *testing.T) { require.False(t, executed, "expected dangerous_tool to NOT be executed because it is denied by inherited permissions") } +type providerSequence struct { + id string + mu sync.Mutex + streams []chat.MessageStream + calls int +} + +func (p *providerSequence) ID() modelsdev.ID { return modelsdev.ParseIDOrZero(p.id) } + +func (p *providerSequence) CreateChatCompletionStream(context.Context, []chat.Message, []tools.Tool) (chat.MessageStream, error) { + p.mu.Lock() + defer p.mu.Unlock() + p.calls++ + if len(p.streams) == 0 { + return &mockStream{}, nil + } + stream := p.streams[0] + p.streams = p.streams[1:] + return stream, nil +} + +func (p *providerSequence) BaseConfig() base.Config { return base.Config{} } +func (p *providerSequence) MaxTokens() int { return 0 } + +func (p *providerSequence) callCount() int { + p.mu.Lock() + defer p.mu.Unlock() + return p.calls +} + +func TestIdleStreamRetryAllowance_WholeRunLifecycle(t *testing.T) { + t.Parallel() + + allowance := idleStreamRetryPolicy{enabled: true, parentSessionID: "parent"}.allowance() + assert.True(t, allowance.eligible(errStreamIdle, streamResult{})) + allowance.consume() + assert.False(t, allowance.eligible(errStreamIdle, streamResult{}), "a later turn or fallback cannot receive a second retry") + + started := idleStreamRetryPolicy{enabled: true, parentSessionID: "parent"}.allowance() + assert.False(t, started.eligible(errStreamIdle, streamResult{ResponseStarted: true})) + assert.True(t, started.remaining, "post-response stalls must not consume the allowance") +} + +func TestTransferTask_RetriesIdleStreamOnce(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + stalled := newStalledStream() + primary := &providerSequence{id: "test/delegate", streams: []chat.MessageStream{ + stalled, + newStreamBuilder().AddContent("recovered").AddStopWithUsage(1, 1).Build(), + }} + delegate := agent.New("delegate", "Delegate", + agent.WithModel(primary), + agent.WithHooks(&hooks.Config{ + SubagentStop: []hooks.Hook{{Type: hooks.HookTypeBuiltin, Command: "test_record_transfer_stop"}}, + }), + ) + root := agent.New("root", "Root", + agent.WithModel(primary), + agent.WithHooks(&hooks.Config{ + SubagentStop: []hooks.Hook{{Type: hooks.HookTypeBuiltin, Command: "test_record_transfer_stop"}}, + }), + ) + agent.WithSubAgents(delegate)(root) + + store := session.NewInMemorySessionStore() + rt, err := NewLocalRuntime(t.Context(), team.New(team.WithAgents(root, delegate)), + WithSessionCompaction(false), WithModelStore(mockModelStore{}), WithSessionStore(store)) + require.NoError(t, err) + + recorder := &recordingBuiltin{} + require.NoError(t, rt.hooksRegistry.RegisterBuiltin("test_record_transfer_stop", recorder.hook)) + rt.buildHooksExecutors() + rt.ensureBudget() + + var logBuf bytes.Buffer + previousLogger := slog.Default() + slog.SetDefault(slog.New(slog.NewTextHandler(&logBuf, nil))) + defer slog.SetDefault(previousLogger) + + sess := session.New(session.WithUserMessage("Test")) + require.NoError(t, store.UpdateSession(t.Context(), sess)) + eventCh := make(chan Event, 128) + resultCh := make(chan struct { + result *tools.ToolCallResult + err error + }, 1) + go func() { + result, err := rt.handleTaskTransfer( + t.Context(), sess, transferToolCall("delegate"), NewChannelSink(eventCh), tools.NopRuntime{}, + ) + resultCh <- struct { + result *tools.ToolCallResult + err error + }{result: result, err: err} + }() + + <-stalled.recvStarted + time.Sleep(defaultStreamIdleTimeout - time.Nanosecond) //nolint:forbidigo // Advances synthetic time to the boundary. + assert.Equal(t, 1, primary.callCount(), "retry must not start before the five-minute boundary") + time.Sleep(time.Nanosecond) //nolint:forbidigo // Crosses the synthetic timeout boundary. + synctest.Wait() + outcome := <-resultCh + result, err := outcome.result, outcome.err + + require.NoError(t, err) + require.NotNil(t, result) + assert.Equal(t, "recovered", result.Output) + assert.Equal(t, 2, primary.callCount()) + assert.Equal(t, "root", rt.CurrentAgent().Name(), "transfer must restore the caller") + + warningCount := 0 + completionCount := 0 + var completedChild *session.Session + for len(eventCh) > 0 { + event := <-eventCh + switch typed := event.(type) { + case *WarningEvent: + warningCount++ + case *SubSessionCompletedEvent: + completionCount++ + completedChild, _ = typed.SubSession.(*session.Session) + } + } + assert.Equal(t, 1, warningCount) + assert.Equal(t, 1, completionCount) + assert.Contains(t, logBuf.String(), "parent_session_id="+sess.ID) + assert.Contains(t, logBuf.String(), "reason=stream_idle_before_response") + assert.Contains(t, logBuf.String(), "limit=1") + + stops := recorder.snapshot() + require.Len(t, stops, 1) + assert.Equal(t, "delegate", stops[0].AgentName) + assert.Equal(t, "recovered", stops[0].StopResponse) + + require.NotNil(t, completedChild) + assert.Equal(t, "recovered", completedChild.GetLastAssistantMessageContent()) + }) +} + +func TestRunTurn_BudgetAdmissionErrorStopsLoop(t *testing.T) { + primary := &providerSequence{id: "test/root", streams: []chat.MessageStream{ + newImmediateIdleStream(), + newStreamBuilder().AddContent("must not dispatch").AddStopWithUsage(1, 1).Build(), + }} + root := agent.New("root", "Root", agent.WithModel(primary)) + rt, err := NewLocalRuntime(t.Context(), team.New(team.WithAgents(root)), + WithSessionCompaction(false), WithModelStore(mockModelStore{}), + WithBudget(&latest.BudgetConfig{MaxTokens: 1})) + require.NoError(t, err) + rt.ensureBudget() + + sess := session.New(session.WithUserMessage("Test")) + rt.recordBudget(sess, root, &chat.Usage{InputTokens: 1}, nil, 0, &collectSink{}) + ls := &loopState{idleRetry: &idleStreamRetryAllowance{remaining: true, parentSessionID: "parent"}} + sink := &collectSink{} + + _, span := noop.NewTracerProvider().Tracer("test").Start(t.Context(), "test") + control := rt.runTurn( + t.Context(), sess, root, nil, primary, primary.ID(), 0, span, + nil, ls, sink, + ) + + assert.Equal(t, turnExit, control) + assert.Equal(t, turnEndReasonBudgetExceeded, ls.exitReason) + assert.Equal(t, 1, primary.callCount()) + assert.True(t, ls.idleRetry.remaining) + assert.NotEmpty(t, sink.events) + _, ok := sink.events[0].(*BudgetExceededEvent) + assert.True(t, ok) +} + +func TestRunForwarding_DirectTransferOwnBudgetStopFailsAfterLifecycle(t *testing.T) { + primary := &providerSequence{id: "test/delegate", streams: []chat.MessageStream{ + newStreamBuilder(). + AddToolCallName("call_unknown", "unknown_tool"). + AddToolCallArguments("call_unknown", `{}`). + AddToolCallStopWithUsage(1, 1). + Build(), + newStreamBuilder().AddContent("should not dispatch").AddStopWithUsage(1, 1).Build(), + }} + delegate := agent.New("delegate", "Delegate", agent.WithModel(primary)) + root := agent.New("root", "Root", + agent.WithModel(primary), + agent.WithHooks(&hooks.Config{ + SubagentStop: []hooks.Hook{{Type: hooks.HookTypeBuiltin, Command: "test_record_budget_stop"}}, + }), + ) + agent.WithSubAgents(delegate)(root) + + store := session.NewInMemorySessionStore() + rt, err := NewLocalRuntime(t.Context(), team.New(team.WithAgents(root, delegate)), + WithSessionCompaction(false), WithModelStore(mockModelStore{}), WithSessionStore(store), + WithBudget(&latest.BudgetConfig{MaxTokens: 1})) + require.NoError(t, err) + + recorder := &recordingBuiltin{} + require.NoError(t, rt.hooksRegistry.RegisterBuiltin("test_record_budget_stop", recorder.hook)) + rt.buildHooksExecutors() + rt.ensureBudget() + + parent := session.New(session.WithUserMessage("Test")) + require.NoError(t, store.UpdateSession(t.Context(), parent)) + var events []Event + result, err := rt.runForwarding(t.Context(), parent, EventSinkFunc(func(event Event) { + events = append(events, event) + }), delegationRequest{ + SubSessionConfig: SubSessionConfig{AgentName: "delegate"}, + SwitchCurrentAgent: true, + directNativeTransfer: true, + }) + + require.Error(t, err) + assert.Nil(t, result) + assert.Contains(t, err.Error(), "max_tokens") + assert.Equal(t, 1, primary.callCount()) + + budgetCount := 0 + completionCount := 0 + var completedChild *session.Session + for _, event := range events { + switch typed := event.(type) { + case *BudgetExceededEvent: + budgetCount++ + case *SubSessionCompletedEvent: + completionCount++ + completedChild, _ = typed.SubSession.(*session.Session) + } + } + assert.Equal(t, 1, budgetCount) + assert.Equal(t, 1, completionCount) + + stops := recorder.snapshot() + require.Len(t, stops, 1) + assert.NotEmpty(t, stops[0].StopResponse) + + require.NotNil(t, completedChild) + assert.NotEmpty(t, completedChild.GetLastAssistantMessageContent()) +} + +func TestRunForwarding_IgnoresMismatchedBudgetStop(t *testing.T) { + observer := &mismatchedBudgetObserver{} + primary := &providerSequence{id: "test/delegate", streams: []chat.MessageStream{ + newStreamBuilder(). + AddToolCallName("call_unknown", "unknown_tool"). + AddToolCallArguments("call_unknown", `{}`). + AddToolCallStopWithUsage(1, 1). + Build(), + newStreamBuilder().AddContent("must not dispatch").AddStopWithUsage(1, 1).Build(), + }} + delegate := agent.New("delegate", "Delegate", agent.WithModel(primary)) + root := agent.New("root", "Root", agent.WithModel(primary)) + agent.WithSubAgents(delegate)(root) + + rt, err := NewLocalRuntime(t.Context(), team.New(team.WithAgents(root, delegate)), + WithSessionCompaction(false), WithModelStore(mockModelStore{}), + WithBudget(&latest.BudgetConfig{MaxTokens: 1}), WithEventObserver(observer)) + require.NoError(t, err) + rt.ensureBudget() + parent := session.New(session.WithUserMessage("Test")) + observer.mismatchedSessionID = parent.ID + + var forwardedBudget *BudgetExceededEvent + result, err := rt.runForwarding(t.Context(), parent, EventSinkFunc(func(event Event) { + if budgetEvent, ok := event.(*BudgetExceededEvent); ok { + forwardedBudget = budgetEvent + } + }), delegationRequest{ + SubSessionConfig: SubSessionConfig{AgentName: "delegate"}, + directNativeTransfer: true, + }) + require.NoError(t, err) + require.NotNil(t, result) + assert.Equal(t, 1, primary.callCount()) + require.NotEmpty(t, observer.childSessionID) + require.NotNil(t, forwardedBudget) + assert.Equal(t, parent.ID, forwardedBudget.SessionID) + assert.NotEqual(t, observer.childSessionID, forwardedBudget.SessionID) +} + +type mismatchedBudgetObserver struct { + mismatchedSessionID string + childSessionID string +} + +func (o *mismatchedBudgetObserver) OnRunStart(_ context.Context, sess *session.Session) { + o.childSessionID = sess.ID +} + +func (o *mismatchedBudgetObserver) OnEvent(_ context.Context, _ *session.Session, event Event) { + if budgetEvent, ok := event.(*BudgetExceededEvent); ok { + budgetEvent.SessionID = o.mismatchedSessionID + } +} + +func TestRunForwarding_CancellationCannotReturnStaleSuccess(t *testing.T) { + primary := &providerSequence{id: "test/delegate", streams: []chat.MessageStream{ + newStreamBuilder().AddContent("stale").AddStopWithUsage(1, 1).Build(), + }} + delegate := agent.New("delegate", "Delegate", agent.WithModel(primary)) + root := agent.New("root", "Root", agent.WithModel(primary)) + agent.WithSubAgents(delegate)(root) + + rt, err := NewLocalRuntime(t.Context(), team.New(team.WithAgents(root, delegate)), + WithSessionCompaction(false), WithModelStore(mockModelStore{})) + require.NoError(t, err) + parent := session.New(session.WithUserMessage("Test")) + ctx, cancel := context.WithCancel(t.Context()) + cancel() + + result, err := rt.runForwarding(ctx, parent, EventSinkFunc(func(Event) {}), delegationRequest{ + SubSessionConfig: SubSessionConfig{AgentName: "delegate"}, + }) + require.ErrorIs(t, err, context.Canceled) + assert.Nil(t, result) +} + +func TestTransferTask_SpendsIdleRetryBeforeFallback(t *testing.T) { + tests := []struct { + name string + primaryStreams []chat.MessageStream + fallbackStreams []chat.MessageStream + fallbackRetries int + wantOutput string + wantPrimaryCalls int + wantFallbackCalls int + wantWarnings int + }{ + { + name: "recovers on primary", + primaryStreams: []chat.MessageStream{ + newImmediateIdleStream(), newStreamBuilder().AddContent("recovered").AddStopWithUsage(1, 1).Build(), + }, + wantOutput: "recovered", wantPrimaryCalls: 2, wantWarnings: 1, + }, + { + name: "spends retry budget before fallback", + primaryStreams: []chat.MessageStream{ + newImmediateIdleStream(), newImmediateIdleStream(), + }, + fallbackStreams: []chat.MessageStream{ + newStreamBuilder().AddContent("fallback").AddStopWithUsage(1, 1).Build(), + }, + wantOutput: "fallback", wantPrimaryCalls: 2, wantFallbackCalls: 1, wantWarnings: 1, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var logBuf bytes.Buffer + previousLogger := slog.Default() + slog.SetDefault(slog.New(slog.NewTextHandler(&logBuf, nil))) + t.Cleanup(func() { slog.SetDefault(previousLogger) }) + + primary := &providerSequence{id: "test/delegate", streams: tt.primaryStreams} + fallback := &providerSequence{id: "test/fallback", streams: tt.fallbackStreams} + delegateOpts := []agent.Opt{agent.WithModel(primary)} + if len(tt.fallbackStreams) > 0 { + delegateOpts = append(delegateOpts, agent.WithFallbackModel(fallback), agent.WithFallbackRetries(tt.fallbackRetries)) + } + delegate := agent.New("delegate", "Delegate", delegateOpts...) + root := agent.New("root", "Root", agent.WithModel(primary)) + agent.WithSubAgents(delegate)(root) + + rt, err := NewLocalRuntime(t.Context(), team.New(team.WithAgents(root, delegate)), + WithSessionCompaction(false), WithModelStore(mockModelStore{})) + require.NoError(t, err) + + rt.ensureBudget() + + sess := session.New(session.WithUserMessage("Test")) + eventCh := make(chan Event, 128) + result, err := rt.handleTaskTransfer( + t.Context(), sess, transferToolCall("delegate"), NewChannelSink(eventCh), tools.NopRuntime{}, + ) + + require.NoError(t, err) + require.NotNil(t, result) + assert.Equal(t, tt.wantOutput, result.Output) + assert.Equal(t, tt.wantPrimaryCalls, primary.callCount()) + assert.Equal(t, tt.wantFallbackCalls, fallback.callCount()) + + warningCount := 0 + for len(eventCh) > 0 { + if _, ok := (<-eventCh).(*WarningEvent); ok { + warningCount++ + } + } + assert.Equal(t, tt.wantWarnings, warningCount) + }) + } +} + +type immediateIdleStream struct{} + +func newImmediateIdleStream() chat.MessageStream { return &immediateIdleStream{} } +func (*immediateIdleStream) Recv() (chat.MessageStreamResponse, error) { + return chat.MessageStreamResponse{}, errStreamIdle +} +func (*immediateIdleStream) Close() {} + func TestTransferTask_PropagatesPermissions(t *testing.T) { t.Parallel() @@ -865,6 +1271,49 @@ func TestTransferTask_RejectsDirectCycle(t *testing.T) { assert.Nil(t, firstSubSession(sess), "rejected delegation must not attach a child session") } +func TestTransferTask_NestedTransferGetsFreshIdleRetryAllowance(t *testing.T) { + t.Parallel() + + workerProvider := &providerSequence{id: "test/worker", streams: []chat.MessageStream{ + newImmediateIdleStream(), + newStreamBuilder().AddContent("worker done").AddStopWithUsage(1, 1).Build(), + }} + helperProvider := &providerSequence{id: "test/helper", streams: []chat.MessageStream{ + newImmediateIdleStream(), + newStreamBuilder().AddContent("helper done").AddStopWithUsage(1, 1).Build(), + }} + root := agent.New("root", "Root", agent.WithModel(workerProvider)) + worker := agent.New("worker", "Worker", agent.WithModel(workerProvider)) + helper := agent.New("helper", "Helper", agent.WithModel(helperProvider)) + agent.WithSubAgents(worker)(root) + agent.WithSubAgents(helper)(worker) + + rt, err := NewLocalRuntime(t.Context(), team.New(team.WithAgents(root, worker, helper)), + WithSessionCompaction(false), WithModelStore(mockModelStore{})) + require.NoError(t, err) + rt.ensureBudget() + + parent := session.New(session.WithUserMessage("Test")) + result, err := rt.handleTaskTransfer( + t.Context(), parent, transferToolCall("worker"), EventSinkFunc(func(Event) {}), tools.NopRuntime{}, + ) + require.NoError(t, err) + require.NotNil(t, result) + assert.Equal(t, "worker done", result.Output) + assert.Equal(t, 2, workerProvider.callCount()) + + child := firstSubSession(parent) + require.NotNil(t, child) + child.AgentName = "worker" + result, err = rt.handleTaskTransfer( + t.Context(), child, transferToolCall("helper"), EventSinkFunc(func(Event) {}), tools.NopRuntime{}, + ) + require.NoError(t, err) + require.NotNil(t, result) + assert.Equal(t, "helper done", result.Output) + assert.Equal(t, 2, helperProvider.callCount()) +} + // TestTransferTask_NestedFromPinnedBackgroundSession covers the #3886 nested // scenario: a background agent (pinned session) calling transfer_task. The // caller must resolve from the session, so acyclic multi-level delegation diff --git a/pkg/runtime/fallback.go b/pkg/runtime/fallback.go index fe25990b22..a70ed0fefd 100644 --- a/pkg/runtime/fallback.go +++ b/pkg/runtime/fallback.go @@ -74,7 +74,12 @@ func newFallbackExecutor() *fallbackExecutor { return &fallbackExecutor{} } -// buildModelChain returns the ordered list of models to try: primary first, then fallbacks. +type idleRetryAdmission func() error + +type budgetAdmissionError struct{} + +func (budgetAdmissionError) Error() string { return "delegated idle retry stopped by budget" } + func buildModelChain(primary provider.Provider, fallbacks []provider.Provider) []modelWithFallback { chain := make([]modelWithFallback, 0, 1+len(fallbacks)) chain = append(chain, modelWithFallback{ @@ -182,6 +187,41 @@ func (e *fallbackExecutor) recordSuccess(a *agent.Agent, modelEntry modelWithFal } } +type idleStreamRetryAllowance struct { + remaining bool + parentSessionID string +} + +// idleStreamRetryPolicy is default-disabled. Callers opt in only for direct +// transfer_task children. +type idleStreamRetryPolicy struct { + enabled bool + parentSessionID string +} + +func defaultIdleStreamRetryPolicy() idleStreamRetryPolicy { + return idleStreamRetryPolicy{} +} + +func (p idleStreamRetryPolicy) allowance() *idleStreamRetryAllowance { + return &idleStreamRetryAllowance{ + remaining: p.enabled, + parentSessionID: p.parentSessionID, + } +} + +func isRetryableIdleStream(err error, result streamResult) bool { + return errors.Is(err, errStreamIdle) && !result.ResponseStarted +} + +func (a *idleStreamRetryAllowance) eligible(err error, result streamResult) bool { + return a != nil && a.remaining && isRetryableIdleStream(err, result) +} + +func (a *idleStreamRetryAllowance) consume() { + a.remaining = false +} + // classifyAttemptError handles an error from a stream attempt: checks for // context cancellation, classifies the error, and returns either a // per-iteration decision (retry the same model or skip to the next) or a @@ -214,7 +254,9 @@ func (e *fallbackExecutor) classifyAttemptError( } // execute attempts to create a stream and get a response using the primary model, -// falling back to configured fallback models if the primary fails. +// falling back to configured fallback models if the primary fails. When +// idleRetry has allowance, the first idle timeout before any response payload is +// retried immediately once across the whole delegated child run. // // Retry behavior: // - Retryable errors (5xx, timeouts): retry the same model with exponential backoff @@ -236,6 +278,8 @@ func (e *fallbackExecutor) execute( sess *session.Session, m *modelsdev.Model, events EventSink, + idleRetry *idleStreamRetryAllowance, + admitIdleRetry idleRetryAdmission, ) (streamResult, provider.Provider, error) { fallbackModels := a.FallbackModels() fallbackRetries := getEffectiveRetries(a) @@ -273,7 +317,6 @@ func (e *fallbackExecutor) execute( fbSpan.SetOutcome(genai.FallbackOutcomeContextCanceled) return streamResult{}, nil, ctx.Err() } - fbSpan.IncrementAttempt() // Apply backoff before retry (not on first attempt of each model) if attempt > 0 { @@ -315,6 +358,8 @@ func (e *fallbackExecutor) execute( // the goroutine reading the response body. streamCtx, streamCancel := context.WithCancelCause(ctx) + // Count only requests that reach the provider dispatch boundary. + fbSpan.IncrementAttempt() stream, err := modelEntry.provider.CreateChatCompletionStream(streamCtx, attemptMessages, agentTools) if err != nil { streamCancel(nil) @@ -346,6 +391,51 @@ func (e *fallbackExecutor) execute( streamCancel(nil) // always release the child context if err != nil { lastErr = err + if idleRetry.eligible(err, res) && admitIdleRetry != nil { + if ctx.Err() != nil { + fbSpan.SetOutcome(genai.FallbackOutcomeContextCanceled) + return streamResult{}, nil, ctx.Err() + } + if err := admitIdleRetry(); err != nil { + fbSpan.SetOutcome(genai.FallbackOutcomeFailed) + return streamResult{}, nil, err + } + if ctx.Err() != nil { + fbSpan.SetOutcome(genai.FallbackOutcomeContextCanceled) + return streamResult{}, nil, ctx.Err() + } + idleRetry.consume() + modelID := modelEntry.provider.ID().String() + const retryLimit = 1 + slog.WarnContext(ctx, "Delegated model stream idle before response; retrying immediately", + "agent", a.Name(), + "model", modelID, + "session_id", sess.ID, + "parent_session_id", idleRetry.parentSessionID, + "reason", "stream_idle_before_response", + "retry", retryLimit, + "limit", retryLimit, + ) + events.Emit(Warning( + "Delegated model stream was idle before responding; retrying once.", + a.Name(), + )) + + streamCtx, streamCancel = context.WithCancelCause(ctx) + fbSpan.IncrementAttempt() + stream, err = modelEntry.provider.CreateChatCompletionStream(streamCtx, attemptMessages, agentTools) + if err == nil { + res, err = handleStream(streamCtx, streamCancel, stream, a, agentTools, sess, m, e.telemetry, events, defaultStreamIdleTimeout) + } + streamCancel(nil) + if err == nil { + e.recordSuccess(a, modelEntry, primaryFailedWithNonRetryable) + fbSpan.SetFinalModel(modelEntry.provider.ID().Model) + fbSpan.SetOutcome(genai.FallbackOutcomeSuccess) + return res, modelEntry.provider, nil + } + lastErr = err + } decision, retErr := e.classifyAttemptError(ctx, err, a, modelEntry, attempt, hasFallbacks, &primaryFailedWithNonRetryable) if retErr != nil { fbSpan.SetOutcome(genai.FallbackOutcomeContextCanceled) diff --git a/pkg/runtime/fallback_test.go b/pkg/runtime/fallback_test.go index ac4c320737..2b1fd4360d 100644 --- a/pkg/runtime/fallback_test.go +++ b/pkg/runtime/fallback_test.go @@ -3,6 +3,8 @@ package runtime import ( "context" "errors" + "fmt" + "sync" "testing" "testing/synctest" "time" @@ -95,6 +97,29 @@ func TestBuildModelChain(t *testing.T) { }) } +func TestIsRetryableIdleStream(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + err error + result streamResult + want bool + }{ + {name: "sentinel before response", err: fmt.Errorf("wrapped: %w", errStreamIdle), want: true}, + {name: "sentinel after response", err: errStreamIdle, result: streamResult{ResponseStarted: true}}, + {name: "matching text without sentinel", err: errors.New("model stream stalled after 30s with no data")}, + {name: "other error", err: errors.New("connection reset")}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + assert.Equal(t, tt.want, isRetryableIdleStream(tt.err, tt.result)) + }) + } +} + func TestFallbackOrder(t *testing.T) { t.Parallel() @@ -673,3 +698,198 @@ func TestRateLimitGate_EnabledWithFallbacks_SkipsToFallback(t *testing.T) { assert.Equal(t, 1, primary.callCount, "primary should only be called once — fallbacks take priority over retry") }) } + +type idleErrorStream struct{} + +func (*idleErrorStream) Recv() (chat.MessageStreamResponse, error) { + return chat.MessageStreamResponse{}, errStreamIdle +} +func (*idleErrorStream) Close() {} + +func TestFallbackExecutor_IdleWithoutAdmissionUsesOrdinaryClassification(t *testing.T) { + t.Parallel() + + modelProvider := &mockProvider{id: "test/model", stream: &idleErrorStream{}} + a := agent.New("delegate", "Delegate", agent.WithModel(modelProvider)) + executor := newFallbackExecutor() + executor.cooldowns = newCooldownManager(time.Now) + executor.telemetry = defaultTelemetry{} + + _, _, err := executor.execute( + t.Context(), a, modelProvider, nil, nil, session.New(), nil, + &collectSink{}, &idleStreamRetryAllowance{remaining: true, parentSessionID: "parent"}, nil, + ) + + require.Error(t, err) + require.ErrorIs(t, err, errStreamIdle) +} + +type idleThenResultProvider struct { + id string + mu sync.Mutex + streams []chat.MessageStream + calls int +} + +func (p *idleThenResultProvider) ID() modelsdev.ID { return modelsdev.ParseIDOrZero(p.id) } +func (p *idleThenResultProvider) CreateChatCompletionStream(context.Context, []chat.Message, []tools.Tool) (chat.MessageStream, error) { + p.mu.Lock() + defer p.mu.Unlock() + p.calls++ + stream := p.streams[0] + p.streams = p.streams[1:] + return stream, nil +} +func (p *idleThenResultProvider) BaseConfig() base.Config { return base.Config{} } +func (p *idleThenResultProvider) MaxTokens() int { return 0 } +func (p *idleThenResultProvider) callCount() int { + p.mu.Lock() + defer p.mu.Unlock() + return p.calls +} + +func TestFallbackExecutor_IdleRetryIsAdditionalImmediateRequest(t *testing.T) { + t.Parallel() + + recovered := newStreamBuilder().AddContent("recovered").AddStopWithUsage(1, 1).Build() + modelProvider := &idleThenResultProvider{id: "test/model", streams: []chat.MessageStream{&idleErrorStream{}, recovered}} + a := agent.New("delegate", "Delegate", agent.WithModel(modelProvider), agent.WithFallbackRetries(-1)) + executor := newFallbackExecutor() + executor.cooldowns = newCooldownManager(time.Now) + executor.telemetry = defaultTelemetry{} + sink := &collectSink{} + allowance := &idleStreamRetryAllowance{remaining: true, parentSessionID: "parent"} + result, used, err := executor.execute( + t.Context(), a, modelProvider, nil, nil, session.New(), nil, sink, allowance, + func() error { return nil }, + ) + + require.NoError(t, err) + assert.Equal(t, "recovered", result.Content) + assert.Equal(t, modelProvider, used) + assert.Equal(t, 2, modelProvider.callCount(), "maxAttempts=1 must still admit one additional request") + assert.False(t, allowance.remaining) + warnings := 0 + fallbacks := 0 + for _, event := range sink.events { + switch event.(type) { + case *WarningEvent: + warnings++ + case *ModelFallbackEvent: + fallbacks++ + } + } + assert.Equal(t, 1, warnings) + assert.Zero(t, fallbacks) +} + +func TestFallbackExecutor_DeniedIdleRetryPreservesAllowanceAndEmitsNothing(t *testing.T) { + t.Parallel() + + modelProvider := &idleThenResultProvider{id: "test/model", streams: []chat.MessageStream{&idleErrorStream{}}} + a := agent.New("delegate", "Delegate", agent.WithModel(modelProvider), agent.WithFallbackRetries(-1)) + executor := newFallbackExecutor() + executor.cooldowns = newCooldownManager(time.Now) + executor.telemetry = defaultTelemetry{} + sink := &collectSink{} + allowance := &idleStreamRetryAllowance{remaining: true, parentSessionID: "parent"} + admissionErr := errors.New("not admitted") + + _, _, err := executor.execute( + t.Context(), a, modelProvider, nil, nil, session.New(), nil, sink, allowance, + func() error { return admissionErr }, + ) + + require.ErrorIs(t, err, admissionErr) + assert.Equal(t, 1, modelProvider.callCount()) + assert.True(t, allowance.remaining) + assert.Empty(t, sink.events) +} + +type cancelOnRecvIdleStream struct { + cancel context.CancelFunc +} + +func (s *cancelOnRecvIdleStream) Recv() (chat.MessageStreamResponse, error) { + s.cancel() + return chat.MessageStreamResponse{}, errStreamIdle +} + +func (*cancelOnRecvIdleStream) Close() {} + +func TestFallbackExecutor_CancellationBeforeIdleRetryAdmission(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(t.Context()) + modelProvider := &mockProvider{id: "test/model", stream: &cancelOnRecvIdleStream{cancel: cancel}} + a := agent.New("delegate", "Delegate", agent.WithModel(modelProvider)) + executor := newFallbackExecutor() + executor.cooldowns = newCooldownManager(time.Now) + executor.telemetry = defaultTelemetry{} + admissionChecks := 0 + + _, _, err := executor.execute( + ctx, a, modelProvider, nil, nil, session.New(), nil, + &collectSink{}, &idleStreamRetryAllowance{remaining: true, parentSessionID: "parent"}, + func() error { + admissionChecks++ + return nil + }, + ) + + require.ErrorIs(t, err, context.Canceled) + assert.Zero(t, admissionChecks) +} + +func TestFallbackExecutor_CancellationAfterIdleRetryAdmission(t *testing.T) { + t.Parallel() + + ctx, cancel := context.WithCancel(t.Context()) + modelProvider := &idleThenResultProvider{id: "test/model", streams: []chat.MessageStream{&idleErrorStream{}}} + a := agent.New("delegate", "Delegate", agent.WithModel(modelProvider)) + executor := newFallbackExecutor() + executor.cooldowns = newCooldownManager(time.Now) + executor.telemetry = defaultTelemetry{} + sink := &collectSink{} + allowance := &idleStreamRetryAllowance{remaining: true, parentSessionID: "parent"} + admissionChecks := 0 + + _, _, err := executor.execute( + ctx, a, modelProvider, nil, nil, session.New(), nil, sink, allowance, + func() error { + admissionChecks++ + cancel() + return nil + }, + ) + + require.ErrorIs(t, err, context.Canceled) + assert.Equal(t, 1, admissionChecks) + assert.Equal(t, 1, modelProvider.callCount()) + assert.True(t, allowance.remaining) + assert.Empty(t, sink.events) +} + +func TestFallbackExecutor_IdleRetryStopsForBudget(t *testing.T) { + t.Parallel() + + modelProvider := &mockProvider{id: "test/model", stream: &idleErrorStream{}} + a := agent.New("delegate", "Delegate", agent.WithModel(modelProvider)) + executor := newFallbackExecutor() + executor.cooldowns = newCooldownManager(time.Now) + executor.telemetry = defaultTelemetry{} + admissionChecks := 0 + admissionErr := errors.New("budget exceeded") + + _, _, err := executor.execute( + t.Context(), a, modelProvider, nil, nil, session.New(), nil, + &collectSink{}, &idleStreamRetryAllowance{remaining: true, parentSessionID: "parent"}, + func() error { + admissionChecks++ + return admissionErr + }, + ) + + require.ErrorIs(t, err, admissionErr) + assert.Equal(t, 1, admissionChecks) +} diff --git a/pkg/runtime/live_sessions.go b/pkg/runtime/live_sessions.go index eea91b4708..ac924a4beb 100644 --- a/pkg/runtime/live_sessions.go +++ b/pkg/runtime/live_sessions.go @@ -75,6 +75,9 @@ type liveSessionEntry struct { // foreground parent finishes and unregisters. Immutable once the entry // is published. treeRootID string + // idleRetry owns the single pre-response idle retry allowance for the + // complete delegated transfer_task child run. + idleRetry *idleStreamRetryAllowance // compactCh holds at most one pending explicit compaction request. // Sends happen only under liveSessionsMu while the entry is still // registered; the final drain runs after unregistration, so an @@ -86,9 +89,14 @@ type liveSessionEntry struct { // RunStream before the run goroutine starts so the session is targetable for // the whole lifetime of its stream. func (r *LocalRuntime) registerLiveSession(sess *session.Session) *liveSessionEntry { + return r.registerLiveSessionWithIdleRetry(sess, defaultIdleStreamRetryPolicy()) +} + +func (r *LocalRuntime) registerLiveSessionWithIdleRetry(sess *session.Session, policy idleStreamRetryPolicy) *liveSessionEntry { entry := &liveSessionEntry{ sess: sess, agentName: r.sessionAgentName(sess), + idleRetry: policy.allowance(), compactCh: make(chan liveCompactionRequest, 1), } r.liveSessionsMu.Lock() diff --git a/pkg/runtime/loop.go b/pkg/runtime/loop.go index b9fb56a5e4..1f3571013e 100644 --- a/pkg/runtime/loop.go +++ b/pkg/runtime/loop.go @@ -2,6 +2,7 @@ package runtime import ( "context" + "errors" "fmt" "log/slog" "path" @@ -247,6 +248,10 @@ func (r *LocalRuntime) streamStoppedTimeout() time.Duration { // the response, executes any tool calls, and loops until the model signals stop // or the iteration limit is reached. func (r *LocalRuntime) RunStream(ctx context.Context, sess *session.Session) <-chan Event { + return r.runStream(ctx, sess, defaultIdleStreamRetryPolicy()) +} + +func (r *LocalRuntime) runStream(ctx context.Context, sess *session.Session, idleRetryPolicy idleStreamRetryPolicy) <-chan Event { slog.DebugContext(ctx, "Starting runtime stream", "agent", r.currentAgentName(), "session_id", sess.ID) events := make(chan Event, defaultEventChannelCapacity) rootStream := !sess.IsSubSession() @@ -261,7 +266,7 @@ func (r *LocalRuntime) RunStream(ctx context.Context, sess *session.Session) <-c // Register before the run goroutine starts so the session is listed in // the /context team view (and targetable for explicit compaction) for // the whole lifetime of its stream. - entry := r.registerLiveSession(sess) + entry := r.registerLiveSessionWithIdleRetry(sess, idleRetryPolicy) go func() { if rootStream { @@ -374,6 +379,7 @@ func (r *LocalRuntime) runStreamLoop(ctx context.Context, sess *session.Session, sessionStart := r.executeSessionStartHooks(ctx, sess, a, sink) ls := &loopState{ maxIterations: sess.MaxIterations, + idleRetry: liveEntry.idleRetry, sessionStartMsgs: sessionStart.messages, sessionStartLegacyMsgs: sessionStart.legacyMessages(), sessionStartSources: sessionStart.sources, @@ -661,6 +667,8 @@ type loopState struct { // empty-response warning would otherwise imply. Reset on agent switch // so it never carries across agents. prevTurnMadeToolCalls bool + // idleRetry is shared by every turn and fallback attempt in this child run. + idleRetry *idleStreamRetryAllowance } // emptyTurnWarning classifies an empty assistant turn (no content, no tool @@ -813,9 +821,23 @@ func (r *LocalRuntime) runTurn( // Runtime message transforms run inside fallback.execute so each attempt // uses the capabilities of the provider that will receive it. - // Try primary model with fallback chain if configured + // Try primary model with fallback chain if configured. The idle retry + // admission callback belongs to this loop invocation so a denial can stop + // this exact session before generic model-error handling runs. agentTools = r.toolDeferrals.MarkAt(sess.ID, lastToolCallID(messages), agentTools) - res, usedModel, err := r.fallback.execute(streamCtx, a, model, messages, agentTools, sess, m, events) + admitIdleRetry := func() error { + if r.enforceBudget(ctx, sess, a, events) == iterationStop { + return budgetAdmissionError{} + } + return nil + } + res, usedModel, err := r.fallback.execute(streamCtx, a, model, messages, agentTools, sess, m, events, ls.idleRetry, admitIdleRetry) + var budgetStop budgetAdmissionError + if errors.As(err, &budgetStop) { + endStreamSpan() + endReason = turnEndReasonBudgetExceeded + return turnExit + } if err != nil { outcome := r.handleStreamError(ctx, sess, a, err, contextLimit, &ls.overflowCompactions, streamSpan, events) endStreamSpan() diff --git a/pkg/runtime/streaming.go b/pkg/runtime/streaming.go index 00d934baa2..a5c53cf2d1 100644 --- a/pkg/runtime/streaming.go +++ b/pkg/runtime/streaming.go @@ -47,6 +47,7 @@ type streamResult struct { ReasoningContent string ThinkingSignature string ThoughtSignature []byte + ResponseStarted bool // Media accumulates every [chat.MediaDelta] streamed during the turn // (e.g. generated images). Populated regardless of provider — see // chat.MessageDelta.Media. @@ -117,6 +118,11 @@ func handleStream(ctx context.Context, cancelStream context.CancelCauseFunc, str var media []chat.MediaDelta var messageUsage *chat.Usage var providerFinishReason chat.FinishReason + var responseStarted bool + + failedResult := func() streamResult { + return streamResult{Stopped: true, ResponseStarted: responseStarted} + } toolCallIndex := make(map[string]int) // toolCallID -> index in toolCalls slice emittedPartial := make(map[string]bool) // toolCallID -> whether we've emitted a partial event @@ -225,7 +231,7 @@ mainLoop: break mainLoop } if res.err != nil { - return streamResult{Stopped: true}, fmt.Errorf("error receiving from stream: %w", res.err) + return failedResult(), fmt.Errorf("error receiving from stream: %w", res.err) } response := res.response @@ -243,11 +249,13 @@ mainLoop: choice := response.Choices[0] if len(choice.Delta.ThoughtSignature) > 0 { + responseStarted = true thoughtSignature = choice.Delta.ThoughtSignature } // A terminal chunk can also carry media; collect it before returning. if len(choice.Delta.Media) > 0 { + responseStarted = true media = append(media, choice.Delta.Media...) } @@ -258,6 +266,7 @@ mainLoop: // reason first would drop the call and the turn would end with an // empty assistant message ("No response from agent"). if len(choice.Delta.ToolCalls) > 0 { + responseStarted = true // Process each tool call delta for _, delta := range choice.Delta.ToolCalls { idx, exists := toolCallIndex[delta.ID] @@ -347,6 +356,7 @@ mainLoop: Stopped: len(toolCalls) == 0, // stop only when there are no tool calls to execute FinishReason: finishReason, Usage: messageUsage, + ResponseStarted: responseStarted, }, nil } @@ -358,23 +368,26 @@ mainLoop: } if choice.Delta.ReasoningContent != "" { + responseStarted = true events.Emit(AgentChoiceReasoning(a.Name(), sess.ID, choice.Delta.ReasoningContent)) fullReasoningContent.WriteString(choice.Delta.ReasoningContent) } // Capture thinking signature for Anthropic extended thinking if choice.Delta.ThinkingSignature != "" { + responseStarted = true thinkingSignature = choice.Delta.ThinkingSignature } if choice.Delta.Content != "" { + responseStarted = true appendContent(markerFilter.Push(choice.Delta.Content)) } case <-ctx.Done(): // Context cancelled (SIGTERM, Ctrl+C, or idle-timeout cancel from // this function). Return promptly so graceful shutdown can proceed. - return streamResult{Stopped: true}, ctx.Err() + return failedResult(), ctx.Err() case <-idleTimer.C: slog.WarnContext(ctx, "Model stream stalled: no data received within idle timeout", @@ -387,7 +400,7 @@ mainLoop: if cancelStream != nil { cancelStream(errStreamIdle) } - return streamResult{Stopped: true}, fmt.Errorf("model stream stalled after %s with no data: %w", + return failedResult(), fmt.Errorf("model stream stalled after %s with no data: %w", idleTimeout, errStreamIdle) } } @@ -441,5 +454,6 @@ mainLoop: Stopped: stoppedNoToolCalls, FinishReason: finishReason, Usage: messageUsage, + ResponseStarted: responseStarted, }, nil } diff --git a/pkg/runtime/streaming_test.go b/pkg/runtime/streaming_test.go index 4925e682a0..5990286a2e 100644 --- a/pkg/runtime/streaming_test.go +++ b/pkg/runtime/streaming_test.go @@ -6,6 +6,7 @@ import ( "strings" "sync" "testing" + "testing/synctest" "time" "github.com/stretchr/testify/assert" @@ -441,29 +442,128 @@ func (s *stalledStream) Close() { // wrapping errStreamIdle when no SSE chunk arrives within the idle window. // It also checks that the provided cancelStream function is called so the // HTTP transport can close the underlying TCP connection. +func TestHandleStream_ProductionIdleTimeoutBoundary(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + stream := newStalledStream() + a := agent.New("root", "test", agent.WithModel(&mockProvider{id: "test/mock-model", stream: stream})) + sess := session.New(session.WithUserMessage("go")) + resultCh := make(chan error, 1) + + go func() { + _, err := handleStream( + t.Context(), func(error) { stream.Close() }, stream, a, nil, sess, nil, + defaultTelemetry{}, NewChannelSink(make(chan Event, 64)), defaultStreamIdleTimeout, + ) + resultCh <- err + }() + + <-stream.recvStarted + time.Sleep(defaultStreamIdleTimeout - time.Nanosecond) //nolint:forbidigo // Advances synthetic time to the boundary. + select { + case err := <-resultCh: + t.Fatalf("production idle timeout fired before five minutes: %v", err) + default: + } + time.Sleep(time.Nanosecond) //nolint:forbidigo // Crosses the synthetic timeout boundary. + synctest.Wait() + select { + case err := <-resultCh: + require.ErrorIs(t, err, errStreamIdle) + default: + t.Fatal("production idle timeout did not fire at five minutes") + } + }) +} + func TestHandleStream_IdleTimeout(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + stream := newStalledStream() + a := agent.New("root", "test", agent.WithModel(&mockProvider{id: "test/mock-model", stream: stream})) + sess := session.New(session.WithUserMessage("go")) + + cancelCalled := false + cancelStream := func(cause error) { + cancelCalled = true + stream.Close() // unblock the stalled Recv so the reader goroutine can exit + } + + evCh := make(chan Event, 64) + res, err := handleStream( + t.Context(), cancelStream, stream, a, nil, sess, nil, + defaultTelemetry{}, NewChannelSink(evCh), 30*time.Second, + ) + + require.Error(t, err) + require.ErrorIs(t, err, errStreamIdle, "error must wrap errStreamIdle") + assert.True(t, res.Stopped) + assert.False(t, res.ResponseStarted) + assert.True(t, cancelCalled, "cancelStream must be called on idle timeout") + }) +} + +func TestHandleStream_IdleTimeoutAfterResponseStarted(t *testing.T) { t.Parallel() - stream := newStalledStream() - a := agent.New("root", "test", agent.WithModel(&mockProvider{id: "test/mock-model", stream: stream})) - sess := session.New(session.WithUserMessage("go")) + tests := []struct { + name string + response chat.MessageStreamResponse + }{ + {name: "content", response: chat.MessageStreamResponse{Choices: []chat.MessageStreamChoice{{Delta: chat.MessageDelta{Content: "partial"}}}}}, + {name: "reasoning", response: chat.MessageStreamResponse{Choices: []chat.MessageStreamChoice{{Delta: chat.MessageDelta{ReasoningContent: "thinking"}}}}}, + {name: "thinking signature", response: chat.MessageStreamResponse{Choices: []chat.MessageStreamChoice{{Delta: chat.MessageDelta{ThinkingSignature: "signature"}}}}}, + {name: "thought signature", response: chat.MessageStreamResponse{Choices: []chat.MessageStreamChoice{{Delta: chat.MessageDelta{ThoughtSignature: []byte("signature")}}}}}, + {name: "tool call", response: streamResponseWithToolCall()}, + {name: "media", response: chat.MessageStreamResponse{Choices: []chat.MessageStreamChoice{{Delta: chat.MessageDelta{Media: []chat.MediaDelta{{}}}}}}}, + } - cancelCalled := false - cancelStream := func(cause error) { - cancelCalled = true - stream.Close() // unblock the stalled Recv so the reader goroutine can exit + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + stream := newResponseThenStalledStream(tt.response) + a := agent.New("root", "test", agent.WithModel(&mockProvider{id: "test/mock-model", stream: stream})) + sess := session.New(session.WithUserMessage("go")) + cancelStream := func(error) { stream.Close() } + + res, err := handleStream( + t.Context(), cancelStream, stream, a, nil, sess, nil, + defaultTelemetry{}, NewChannelSink(make(chan Event, 64)), 50*time.Millisecond, + ) + + require.ErrorIs(t, err, errStreamIdle) + assert.True(t, res.ResponseStarted) + }) } +} - evCh := make(chan Event, 64) - res, err := handleStream( - t.Context(), cancelStream, stream, a, nil, sess, nil, - defaultTelemetry{}, NewChannelSink(evCh), 50*time.Millisecond, - ) +func streamResponseWithToolCall() chat.MessageStreamResponse { + return chat.MessageStreamResponse{Choices: []chat.MessageStreamChoice{{ + Delta: chat.MessageDelta{ToolCalls: []tools.ToolCall{{ + ID: "call", Function: tools.FunctionCall{Name: "tool"}, + }}}, + }}} +} - require.Error(t, err) - require.ErrorIs(t, err, errStreamIdle, "error must wrap errStreamIdle") - assert.True(t, res.Stopped) - assert.True(t, cancelCalled, "cancelStream must be called on idle timeout") +type responseThenStalledStream struct { + response chat.MessageStreamResponse + stalled *stalledStream + sent bool +} + +func newResponseThenStalledStream(response chat.MessageStreamResponse) *responseThenStalledStream { + return &responseThenStalledStream{response: response, stalled: newStalledStream()} +} + +func (s *responseThenStalledStream) Recv() (chat.MessageStreamResponse, error) { + if !s.sent { + s.sent = true + return s.response, nil + } + return s.stalled.Recv() +} + +func (s *responseThenStalledStream) Close() { + s.stalled.Close() } // TestHandleStream_ContextCancellation verifies that handleStream returns diff --git a/pkg/teamloader/teamloader.go b/pkg/teamloader/teamloader.go index 5e55f4a92f..b598a74bb3 100644 --- a/pkg/teamloader/teamloader.go +++ b/pkg/teamloader/teamloader.go @@ -402,6 +402,8 @@ func LoadWithConfig(ctx context.Context, agentSource config.Source, runConfig *c } promptFiles = unique + // Build options in source order: author configuration first, followed by + // loader/runtime additions below, so explicit execution overrides win. opts := []agent.Opt{ agent.WithName(agentConfig.Name), agent.WithDescription(expander.Expand(ctx, agentConfig.Description, nil)), @@ -660,6 +662,7 @@ func getModelsForAgent(ctx context.Context, cfg *latest.Config, a *latest.AgentC isAutoModel = true } modelCfg.Name = name + config.ApplyModelOverridePolicy(cfg, a.Name, name, &modelCfg) // Use max_tokens from config if specified, otherwise look up from models.dev maxTokens := &defaultMaxTokens diff --git a/pkg/teamloader/teamloader_test.go b/pkg/teamloader/teamloader_test.go index ab5500a361..d20500f43c 100644 --- a/pkg/teamloader/teamloader_test.go +++ b/pkg/teamloader/teamloader_test.go @@ -5,11 +5,13 @@ import ( "encoding/json" "errors" "io/fs" + "maps" "net/http" "net/http/httptest" "os" "path/filepath" "runtime" + "slices" "sync" "sync/atomic" "testing" @@ -1092,6 +1094,95 @@ func TestLoadWithConfig_WithWorkingDirDoesNotLeak(t *testing.T) { assert.Equal(t, callerDir, runConfig.WorkingDir) } +func TestLoadWithConfigModelOverridePolicyMatrix(t *testing.T) { + t.Setenv("OPENAI_API_KEY", "dummy") + + data := []byte(`models: + serial: + provider: openai + model: original-serial + parallel_tool_calls: false + parallel: + provider: openai + model: original-parallel + parallel_tool_calls: true + unset: + provider: openai + model: original-unset + named_nil: + provider: openai + model: named +agents: + serial: + model: serial + instruction: test + parallel: + model: parallel + instruction: test + unset: + model: unset + instruction: test + named: + model: serial + instruction: test +`) + + result, err := LoadWithConfig( + t.Context(), + config.NewBytesSource("overrides.yaml", data), + &config.RuntimeConfig{}, + withTestProviderRegistry(WithModelOverrides([]string{ + "serial=openai/replacement", + "parallel=openai/replacement", + "unset=openai/replacement", + "named=named_nil", + }))..., + ) + require.NoError(t, err) + + assert.Equal(t, map[string]string{ + "serial": "openai/replacement", "parallel": "openai/replacement", + "unset": "openai/replacement", "named": "named_nil", + }, result.AgentDefaultModels) + assert.ElementsMatch(t, []string{"serial", "parallel", "unset", "named_nil", "openai/replacement"}, slices.Collect(maps.Keys(result.Models))) + for name := range result.Models { + assert.NotContains(t, name, "__cli_model_") + } + + wantPolicy := map[string]*bool{ + "serial": new(false), "parallel": new(true), "unset": nil, "named": nil, + } + wantRef := map[string]string{ + "serial": "openai/replacement", "parallel": "openai/replacement", + "unset": "openai/replacement", "named": "named_nil", + } + for agentName, modelRef := range wantRef { + teamCfg, ok := result.Team.AgentConfig(agentName) + require.True(t, ok) + assert.Equal(t, modelRef, teamCfg.Model) + + a, err := result.Team.Agent(agentName) + require.NoError(t, err) + providers := a.ConfiguredModels() + require.Len(t, providers, 1) + assert.Equal(t, modelRef, providers[0].BaseConfig().ModelConfig.Name) + assert.Equal(t, map[string]string{ + "serial": "openai/replacement", "parallel": "openai/replacement", + "unset": "openai/replacement", "named": "openai/named", + }[agentName], providers[0].ID().String()) + + got := providers[0].BaseConfig().ModelConfig.ParallelToolCalls + want := wantPolicy[agentName] + if want == nil { + assert.Nil(t, got) + } else if assert.NotNil(t, got) { + assert.Equal(t, *want, *got) + } + } + assert.Nil(t, result.Models["openai/replacement"].ParallelToolCalls) + assert.Nil(t, result.Models["named_nil"].ParallelToolCalls) +} + // TestLoadRetainsAgentConfig verifies the loader retains the raw resolved // per-agent config on the team (team.WithAgentConfigs) so the agent inspector // can surface declared toolset allow-lists, limits and flags. It uses a @@ -1200,6 +1291,76 @@ func TestLoadRejectsUncompilableToolModeSchema(t *testing.T) { require.ErrorContains(t, err, "agent root: structured_output") } +func TestLoadPreservesModelPolicyDefaults(t *testing.T) { + t.Setenv("OPENAI_API_KEY", "dummy") + + tests := []struct { + name string + yaml string + want *bool + }{ + { + name: "omitted", + yaml: `models: + configured: + provider: openai + model: gpt-4o +agents: + root: + model: configured + instruction: test +`, + want: nil, + }, + { + name: "explicit true", + yaml: `models: + configured: + provider: openai + model: gpt-4o + parallel_tool_calls: true +agents: + root: + model: configured + instruction: test +`, + want: new(true), + }, + { + name: "explicit false", + yaml: `models: + configured: + provider: openai + model: gpt-4o + parallel_tool_calls: false +agents: + root: + model: configured + instruction: test +`, + want: new(false), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + loaded, err := Load(t.Context(), config.NewBytesSource("model.yaml", []byte(tt.yaml)), + &config.RuntimeConfig{}, withTestProviderRegistry()...) + require.NoError(t, err) + root, err := loaded.Agent("root") + require.NoError(t, err) + models := root.ConfiguredModels() + require.Len(t, models, 1) + got := models[0].BaseConfig().ModelConfig.ParallelToolCalls + if tt.want == nil { + assert.Nil(t, got) + } else if assert.NotNil(t, got) { + assert.Equal(t, *tt.want, *got) + } + }) + } +} + // TestLoadPropagatesSafetyDefaults verifies the author-declared safety // defaults travel from the YAML config to the built team: runtime.safety // lands on the team (team.RuntimeSafety) and agents..safety on the diff --git a/pkg/telemetry/genai/runtime.go b/pkg/telemetry/genai/runtime.go index 3935623382..ccd6c5300d 100644 --- a/pkg/telemetry/genai/runtime.go +++ b/pkg/telemetry/genai/runtime.go @@ -85,8 +85,9 @@ func StartFallback(ctx context.Context, agentName, primaryModel string, inCooldo } } -// IncrementAttempt counts one attempt against the chain. Called once per -// (model × retry) iteration so the final span carries the total count. +// IncrementAttempt counts one provider dispatch against the chain. Call it +// immediately before each CreateChatCompletionStream invocation, including +// the delegated idle-stream redispatch. func (s *FallbackSpan) IncrementAttempt() { if s == nil { return