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
112 changes: 112 additions & 0 deletions agent-flow/agent_flow/agent_runtime.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,112 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

"""Per-role model and backend selection shared by performance workflows."""

from __future__ import annotations

from dataclasses import dataclass
from typing import Any, Mapping

from .config import CLAUDE_CODE_DEFAULT_MODEL, CODEX_DEFAULT_MODEL, BackendKind

AGENTS_FIELD = "agents"
_FIELDS = frozenset({"backend", "model", "reasoning_effort", "extra_mcp_servers"})
_BACKENDS = frozenset({"claude-code", "codex"})


@dataclass(frozen=True)
class AgentConfig:
"""Resolved backend configuration for one workflow role."""

backend: BackendKind
model: str
reasoning_effort: str | None = None
extra_mcp_servers: dict[str, Any] | None = None


def validate_agents(data: Mapping[str, Any], roles: tuple[str, ...]) -> list[str]:
"""Validate the optional ``agents`` block for a workflow's roles."""
if AGENTS_FIELD not in data:
return []

agents = data[AGENTS_FIELD]
if not isinstance(agents, Mapping):
return ["'agents' must be a mapping"]

errors: list[str] = []
unknown = set(agents) - {"defaults", "roles"}
if unknown:
errors.append(f"'agents' has unknown field(s) {sorted(unknown)}")

_validate_block(agents.get("defaults", {}), "agents.defaults", errors)
role_blocks = agents.get("roles", {})
if not isinstance(role_blocks, Mapping):
errors.append("'agents.roles' must be a mapping")
return errors

unknown_roles = set(role_blocks) - set(roles)
if unknown_roles:
errors.append(f"'agents.roles' has unknown role(s) {sorted(unknown_roles)}")
for role, block in role_blocks.items():
_validate_block(block, f"agents.roles.{role}", errors)
return errors


def _validate_block(value: Any, path: str, errors: list[str]) -> None:
if not isinstance(value, Mapping):
errors.append(f"'{path}' must be a mapping")
return

unknown = set(value) - _FIELDS
if unknown:
errors.append(f"'{path}' has unknown field(s) {sorted(unknown)}")
backend = value.get("backend")
if backend is not None and backend not in _BACKENDS:
errors.append(f"'{path}.backend' must be 'claude-code' or 'codex'")
for field in ("model", "reasoning_effort"):
item = value.get(field)
if item is not None and (not isinstance(item, str) or not item.strip()):
errors.append(f"'{path}.{field}' must be a non-empty string")
servers = value.get("extra_mcp_servers")
if servers is not None and not isinstance(servers, Mapping):
errors.append(f"'{path}.extra_mcp_servers' must be a mapping")


def resolve_agent_config(
data: Mapping[str, Any],
role: str,
*,
default_backend: BackendKind = "claude-code",
default_model: str = CLAUDE_CODE_DEFAULT_MODEL,
) -> AgentConfig:
"""Resolve one role from workflow defaults and its optional override."""
agents = data.get(AGENTS_FIELD, {})
if not isinstance(agents, Mapping):
return AgentConfig(default_backend, default_model)
defaults = agents.get("defaults", {})
defaults = defaults if isinstance(defaults, Mapping) else {}
role_blocks = agents.get("roles", {})
role_blocks = role_blocks if isinstance(role_blocks, Mapping) else {}
override = role_blocks.get(role, {})
override = override if isinstance(override, Mapping) else {}

backend = override.get("backend", defaults.get("backend", default_backend))
model = override.get("model")
if model is None:
inherited_model = defaults.get("model")
inherited_backend = defaults.get("backend", default_backend)
if inherited_model is not None and backend == inherited_backend:
model = inherited_model
elif backend == default_backend and "backend" not in defaults and "backend" not in override:
model = default_model
else:
model = CODEX_DEFAULT_MODEL if backend == "codex" else CLAUDE_CODE_DEFAULT_MODEL

servers = override.get("extra_mcp_servers", defaults.get("extra_mcp_servers"))
return AgentConfig(
backend=backend,
model=model,
reasoning_effort=override.get("reasoning_effort", defaults.get("reasoning_effort")),
extra_mcp_servers=dict(servers) if isinstance(servers, Mapping) else None,
)
12 changes: 10 additions & 2 deletions agent-flow/agent_flow/backends/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,11 +14,19 @@

def create_backend(config: BackendConfig | str) -> Backend:
kind = config.kind if isinstance(config, BackendConfig) else config
effort = config.reasoning_effort if isinstance(config, BackendConfig) else None
disabled_skills = config.disabled_skills if isinstance(config, BackendConfig) else ()

if kind == "claude-code":
return claude_code.ClaudeCodeBackend()
return claude_code.ClaudeCodeBackend(
reasoning_effort=effort,
disabled_skills=disabled_skills,
)

if kind == "codex":
return codex.CodexBackend()
return codex.CodexBackend(
reasoning_effort=effort,
disabled_skills=disabled_skills,
)

raise ValueError(f"Unknown backend: {kind!r}")
17 changes: 14 additions & 3 deletions agent-flow/agent_flow/backends/claude_code.py
Original file line number Diff line number Diff line change
Expand Up @@ -446,11 +446,19 @@ async def handler(arguments, definition=definition):


class ClaudeCodeBackend(Backend):
def __init__(
self,
reasoning_effort: str | None = None,
disabled_skills: tuple[str, ...] = (),
) -> None:
self._reasoning_effort = reasoning_effort or _REASONING_EFFORT
self._disabled_skills = disabled_skills

def version(self) -> str:
return _claude_backend_version()

def reasoning_effort(self) -> str:
return _REASONING_EFFORT
return self._reasoning_effort

@asynccontextmanager
async def create_client(
Expand Down Expand Up @@ -490,12 +498,15 @@ async def create_client(
),
mcp_servers=mcp_servers,
model=model,
effort=_REASONING_EFFORT,
effort=self._reasoning_effort,
cwd=cwd or Path.cwd(),
sandbox={"enabled": False},
permission_mode="bypassPermissions",
hooks=hooks,
disallowed_tools=list(disallowed_tools or []),
disallowed_tools=[
*(disallowed_tools or []),
*(f"Skill({name})" for name in self._disabled_skills),
],
)

async with ClaudeSDKClient(options=options) as sdk_client:
Expand Down
29 changes: 25 additions & 4 deletions agent-flow/agent_flow/backends/codex.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from __future__ import annotations

import asyncio
import json
import os
import shutil
import subprocess
Expand Down Expand Up @@ -141,19 +142,36 @@ def _codex_backend_version() -> str:
_REASONING_EFFORT = "max"


def _disabled_skills_override(names: tuple[str, ...]) -> tuple[str, ...]:
if not names:
return ()
entries = ", ".join(f"{{name={json.dumps(name)}, enabled=false}}" for name in names)
return (f"skills.config=[{entries}]",)


class CodexBackend(Backend):
def __init__(self) -> None:
def __init__(
self,
reasoning_effort: str | None = None,
disabled_skills: tuple[str, ...] = (),
) -> None:
self._transport: CodexTransport | None = None
self._reasoning_effort = reasoning_effort or _REASONING_EFFORT
self._disabled_skills = disabled_skills

def version(self) -> str:
return _codex_backend_version()

def reasoning_effort(self) -> str:
return _REASONING_EFFORT
return self._reasoning_effort

async def __aenter__(self) -> "CodexBackend":
self._transport = CodexTransport(
CodexConfig(codex_bin=_resolve_codex_bin(), experimental_api=True)
CodexConfig(
codex_bin=_resolve_codex_bin(),
experimental_api=True,
config_overrides=_disabled_skills_override(self._disabled_skills),
)
)
await self._transport.start()
return self
Expand Down Expand Up @@ -222,7 +240,10 @@ async def create_client(
server["disabled_tools"] = list(
dict.fromkeys([*inherited, *server["disabled_tools"]])
)
config.update(model_reasoning_effort=_REASONING_EFFORT, model_context_window=1000000)
config.update(
model_reasoning_effort=self._reasoning_effort,
model_context_window=1000000,
)
params = ThreadStartParams(
model=model,
developer_instructions=system_prompt or None,
Expand Down
4 changes: 4 additions & 0 deletions agent-flow/agent_flow/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,10 @@
class BackendConfig:
kind: BackendKind
model: str
# Provider reasoning tier. ``None`` retains the backend's historical
# maximum-effort default.
reasoning_effort: str | None = None
disabled_skills: tuple[str, ...] = ()
tools: list[Any] | None = None
# Native SDK hook configuration. Claude accepts HookMatcher callbacks.
# Codex requires hooks to be configured and trusted in its native config;
Expand Down
16 changes: 15 additions & 1 deletion agent-flow/agent_flow/workflows/perf_analyze/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@ benchmarks and profiles it with TensorRT-LLM's
`tensorrt_llm/serve/scripts/benchmark_serving.py`, and writes a report
whose headline is the **main performance bottleneck**.

All roles run on the **Claude Code** backend:
By default, all roles run on the **Claude Code** backend:

```
benchmarker ──▶ projector ──▶ analyzer ──▶ reporter
Expand Down Expand Up @@ -81,6 +81,20 @@ Copy [`task.example.yaml`](./task.example.yaml) and fill it in.

## `task.yaml`

An optional `agents` block selects the backend, model, reasoning effort,
and external MCP servers per role. Unspecified roles retain the historical
Claude defaults:

```yaml
agents:
roles:
projector: {backend: codex, model: gpt-6-astra, reasoning_effort: ultra}
analyzer: {backend: codex, model: gpt-6-astra, reasoning_effort: ultra}
```

Set `casebook.enabled: false` for a control run that hides and blocks the
`perf-optimization-casebook` skill. It is enabled by default.

| Field | Required | Notes |
| --- | --- | --- |
| `checkpoint_path` | ✅ | Model checkpoint dir to serve. Remote when `cluster_ssh` is set; otherwise local. |
Expand Down
12 changes: 10 additions & 2 deletions agent-flow/agent_flow/workflows/perf_analyze/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,11 +4,14 @@
import sys
from pathlib import Path

from agent_flow.agent_runtime import resolve_agent_config

from .prompts import build_perf_analyze_prompts
from .sol_methodology import resolve_sol_methodology
from .state import STATE_FILENAME
from .task_schema import (
TaskSchemaError,
casebook_enabled,
has_slurm_environment,
load_and_validate_task_yaml,
sol_enabled,
Expand Down Expand Up @@ -61,15 +64,19 @@ def _parse_args(argv: list[str] | None = None) -> argparse.Namespace:

def main(argv: list[str] | None = None) -> None:
args = _parse_args(argv)
task_path = args.task
if not args.clean and (args.workspace / STATE_FILENAME).is_file():
task_path = args.workspace / "task.yaml"
try:
task_data = load_and_validate_task_yaml(args.task)
task_data = load_and_validate_task_yaml(task_path)
except TaskSchemaError as exc:
print(f"error: {exc}", file=sys.stderr)
sys.exit(2)
# Resolve the projector's methodology skill once, before the run, so
# it is told to load a skill this session actually has. Skipped (free)
# when the stage is off.
methodology = resolve_sol_methodology(sol_enabled(task_data))
projector_backend = resolve_agent_config(task_data, "projector").backend
methodology = resolve_sol_methodology(sol_enabled(task_data), backend_kind=projector_backend)
note = methodology.console_note()
if note:
print(note, file=sys.stderr)
Expand All @@ -79,6 +86,7 @@ def main(argv: list[str] | None = None) -> None:
sol_methodology=methodology.name,
remote_execution=task_data,
campaign_name=args.workspace.resolve().name,
include_casebook=casebook_enabled(task_data),
)
with PerfAnalyzeWorkflow(
workspace=args.workspace,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@

from ..task_schema import cluster_ssh, remote_run_root
from ._common import (
CASEBOOK_DISABLED,
EXECUTION_SLURM_BOOTSTRAP,
REMOTE_SLURM_EXECUTION,
SOL_ANALYZER_CONTEXT,
Expand Down Expand Up @@ -94,6 +95,7 @@ def build_perf_analyze_prompts(
sol_methodology: str = "full",
remote_execution: Mapping[str, Any] | None = None,
campaign_name: str = "perf-analyze",
include_casebook: bool = True,
) -> PromptBundle:
"""Return the workflow's prompt bundle, optionally augmented.

Expand All @@ -120,6 +122,11 @@ def build_perf_analyze_prompts(
appended to the three roles that may inspect or produce runtime data.
"""
bundle = DEFAULT_PROMPTS
if not include_casebook:
bundle = bundle.with_extensions(
benchmarker=CASEBOOK_DISABLED,
analyzer=CASEBOOK_DISABLED,
)
if sol_methodology != "full":
bundle = dataclasses.replace(bundle, projector=build_projector_prompt(sol_methodology))
if include_slurm_environment:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1085,6 +1085,15 @@
line and proceed — never block the run on it.
"""

CASEBOOK_DISABLED = """\
## Casebook experiment

This run intentionally disables `perf-optimization-casebook`. Do not invoke or
use that skill, even if an earlier section asks for it. Ground conclusions only
in the task, source, measurements, and other enabled skills; `casebook_ref` may
be omitted.
"""


# --------------------------------------------------------------------------- #
# Remote execution boundary (appended with task-specific values by the CLI)
Expand Down
4 changes: 4 additions & 0 deletions agent-flow/agent_flow/workflows/perf_analyze/roles.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,4 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

ROLES = ("benchmarker", "projector", "analyzer", "reporter")
Original file line number Diff line number Diff line change
Expand Up @@ -88,7 +88,9 @@ def console_note(self) -> str | None:
return None


def resolve_sol_methodology(enabled: bool = True) -> SolMethodology:
def resolve_sol_methodology(
enabled: bool = True, backend_kind: str = "claude-code"
) -> SolMethodology:
"""Probe the live skill list and return the methodology the projector has.

Costs one backend connection and no model call (~1 s), and is
Expand All @@ -101,7 +103,9 @@ def resolve_sol_methodology(enabled: bool = True) -> SolMethodology:
from agent_flow.utils import resolve_first_available_skill

try:
loaded, probe_ok = resolve_first_available_skill(SOL_SKILL_CANDIDATES)
loaded, probe_ok = resolve_first_available_skill(
SOL_SKILL_CANDIDATES, backend_kinds=(backend_kind,)
)
except Exception: # noqa: BLE001 - a probe failure must never fail the run
return SolMethodology(probed=False)
if not probe_ok:
Expand Down
Loading
Loading