diff --git a/adapters/claude/fabric-adapter.json b/adapters/claude/fabric-adapter.json index d04516c6..56710f05 100644 --- a/adapters/claude/fabric-adapter.json +++ b/adapters/claude/fabric-adapter.json @@ -41,6 +41,7 @@ "additionalProperties": false }, "config": { + "input": "agent_config", "accepts": [ "models", "models.base_url", diff --git a/adapters/claude/pyproject.toml b/adapters/claude/pyproject.toml index b980d8a9..9f707eec 100644 --- a/adapters/claude/pyproject.toml +++ b/adapters/claude/pyproject.toml @@ -25,6 +25,7 @@ license-files = ["LICENSE"] readme = "pypi.md" requires-python = ">=3.11" dependencies = [ + "nemo-fabric-adapter-contract == 0.2.0", "nemo-fabric-adapters-common == 0.2.0", "tomli-w~=1.2", ] @@ -51,4 +52,5 @@ include = ["nemo_fabric_adapters.claude*"] "share/nemo-fabric/adapters/claude" = ["fabric-adapter.json"] [tool.uv.sources] +nemo-fabric-adapter-contract = { path = "../../adapter-contract", editable = true } nemo-fabric-adapters-common = { path = "../common", editable = true } diff --git a/adapters/claude/src/nemo_fabric_adapters/claude/adapter.py b/adapters/claude/src/nemo_fabric_adapters/claude/adapter.py index 706cc571..7aec44b5 100644 --- a/adapters/claude/src/nemo_fabric_adapters/claude/adapter.py +++ b/adapters/claude/src/nemo_fabric_adapters/claude/adapter.py @@ -30,6 +30,10 @@ from claude_agent_sdk import ResultMessage from claude_agent_sdk import HookMatcher from claude_agent_sdk._errors import MessageParseError +from nemo_fabric_adapter_contract.codec import ContractValidationError +from nemo_fabric_adapter_contract.models import AgentConfig +from nemo_fabric_adapter_contract.models import AgentModelConfig +from nemo_fabric_adapter_contract.models import RuntimeContext from nemo_fabric_adapters.common import lifecycle from nemo_fabric_adapters.common import relay_artifacts from nemo_fabric_adapters.common import relay_gateway @@ -134,16 +138,6 @@ class AdapterRelayError(ClaudeAdapterError): """NeMo Relay setup or lifecycle failure.""" -def _mapping(value: Any, *, name: str) -> dict[str, Any]: - if value is None: - return {} - if not isinstance(value, dict): - raise AdapterConfigError( - "claude_invalid_configuration", f"{name} must be a mapping" - ) - return value - - def _string_list(value: Any, *, name: str) -> list[str]: if value is None: return [] @@ -170,13 +164,14 @@ def _positive_number(value: Any, *, name: str) -> float: return number -def runtime_id(payload: dict[str, Any]) -> str: - value = common_utils.runtime_context(payload).get("runtime_id") - if not isinstance(value, str) or not value: - raise AdapterInputError( - "claude_invalid_request", "NeMo Fabric runtime ID is required" - ) - return value +def _runtime_context(payload: dict[str, Any]) -> RuntimeContext: + try: + return RuntimeContext.from_mapping(payload.get("runtime_context")) + except ContractValidationError as error: + raise lifecycle.LifecycleError( + "claude_invalid_runtime_context", + "Claude runtime context is invalid", + ) from error def request_prompt(payload: dict[str, Any]) -> str: @@ -187,82 +182,58 @@ def request_prompt(payload: dict[str, Any]) -> str: return value -def _settings(payload: dict[str, Any]) -> dict[str, Any]: - return _mapping(common_utils.settings_payload(payload), name="harness.settings") +def _settings(config: AgentConfig) -> dict[str, Any]: + return config.harness.settings if config.harness is not None else {} -def _selected_model_config(payload: dict[str, Any]) -> dict[str, Any]: - return _mapping( - common_utils.selected_model_config(payload), - name="selected model", - ) +def _selected_model_config(config: AgentConfig) -> AgentModelConfig: + model = config.models.get("default") + if model is None and len(config.models) == 1: + model = next(iter(config.models.values())) + if model is None: + raise AdapterConfigError( + "claude_invalid_configuration", + "Claude requires a default model or exactly one model", + ) + return model -def _resolve_path(payload: dict[str, Any], value: str | Path) -> Path: +def _resolve_path(base_dir: str, value: str | Path) -> Path: path = Path(value) if not path.is_absolute(): - path = Path(common_utils.base_dir(payload)) / path + path = Path(base_dir) / path return path -def resolve_cwd(payload: dict[str, Any]) -> Path: - environment = common_utils.environment_payload(payload) - workspace = environment.get("workspace") - return _resolve_path(payload, workspace or common_utils.base_dir(payload)) +def resolve_cwd(runtime_context: RuntimeContext, base_dir: str) -> Path: + return _resolve_path(base_dir, runtime_context.environment.workspace or base_dir) -def selected_model(payload: dict[str, Any]) -> str | None: - model_config = _selected_model_config(payload) - value = model_config.get("model") - if value is None: - return None - provider = model_config.get("provider") - if not isinstance(provider, str) or not provider: - raise AdapterConfigError( - "claude_invalid_configuration", - "selected model provider must be a non-empty string", - ) - if not isinstance(value, str) or not value: - raise AdapterConfigError( - "claude_invalid_configuration", "model must be a non-empty string" - ) - return value.removeprefix("anthropic/") if provider == "anthropic" else value +def selected_model(model: AgentModelConfig) -> str: + return ( + model.model.removeprefix("anthropic/") + if model.provider == "anthropic" + else model.model + ) -def _anthropic_base_url(model: dict[str, Any]) -> str | None: - base_url = common_utils.get_base_url(model) - if base_url is None: +def _anthropic_base_url(model: AgentModelConfig) -> str | None: + if model.base_url is None: return None - if not isinstance(base_url, str) or not base_url: - raise AdapterConfigError( - "claude_invalid_configuration", - "selected model base_url must be a non-empty string", - ) - base_url = base_url.rstrip("/") - return ( - base_url.removesuffix("/v1") - if model.get("provider") != "anthropic" - else base_url - ) + base_url = model.base_url.rstrip("/") + return base_url.removesuffix("/v1") if model.provider != "anthropic" else base_url def _model_environment( - payload: dict[str, Any], environment: dict[str, str] + model: AgentModelConfig, environment: dict[str, str] ) -> dict[str, str]: - model = _selected_model_config(payload) - provider = model.get("provider") - api_key_env = model.get("api_key_env") - if not isinstance(api_key_env, str) or not api_key_env: - if api_key_env is None: - api_key = None - else: - raise AdapterConfigError( - "claude_invalid_configuration", - "selected model api_key_env must be a non-empty string", - ) - else: - api_key = environment.get(api_key_env) or os.environ.get(api_key_env) - if provider != "anthropic" and api_key_env is None: + api_key_env = model.api_key_env + api_key = ( + environment.get(api_key_env) or os.environ.get(api_key_env) + if api_key_env + else None + ) + if model.provider != "anthropic" and api_key_env is None: raise AdapterConfigError( "claude_invalid_configuration", "selected model api_key_env is required for a custom " @@ -274,7 +245,7 @@ def _model_environment( f"{api_key_env} is required for the selected model provider", ) base_url = _anthropic_base_url(model) - if provider != "anthropic" and not base_url: + if model.provider != "anthropic" and not base_url: raise AdapterConfigError( "claude_invalid_configuration", "selected model base_url is required for a custom " @@ -285,35 +256,24 @@ def _model_environment( values["ANTHROPIC_API_KEY"] = api_key if base_url: values["ANTHROPIC_BASE_URL"] = base_url - if provider != "anthropic": + if model.provider != "anthropic": values["ANTHROPIC_AUTH_TOKEN"] = "" return values -def _mcp_servers(payload: dict[str, Any]) -> dict[str, Any]: - native = ( - _mapping(common_utils.capability_plan(payload), name="capability_plan").get( - "native" - ) - or {} - ) - servers = _mapping(native, name="capability_plan.native").get("mcp_servers") or {} +def _mcp_servers(config: AgentConfig) -> dict[str, Any]: + servers = config.mcp.servers if config.mcp is not None else {} result: dict[str, Any] = {} - for name, raw in sorted(_mapping(servers, name="native MCP servers").items()): - server = _mapping(raw, name=f"MCP server {name}") - transport = server.get("transport") - url = server.get("url") - if not isinstance(url, str) or not url: - raise AdapterConfigError( - "claude_invalid_configuration", "MCP server URL is required" - ) + for name, server in sorted(servers.items()): + transport = server.transport + url = server.url if transport == "stdio": result[name] = { "type": "stdio", "command": url, - "args": common_utils.normalize_list(server.get("args")), + "args": server.args, } - if env := server.get("env"): + if env := server.env: result[name]["env"] = env elif transport in {"http", "streamable-http"}: result[name] = {"type": "http", "url": url} @@ -327,7 +287,9 @@ def _mcp_servers(payload: dict[str, Any]) -> dict[str, Any]: return result -def _stage_mcp_config(payload: dict[str, Any]) -> ClaudeMcpSettings | None: +def _stage_mcp_config( + config: AgentConfig, runtime_context: RuntimeContext, base_dir: str +) -> ClaudeMcpSettings | None: # Dictionary-valued ClaudeAgentOptions.mcp_servers are JSON-serialized by # claude-agent-sdk into the literal `--mcp-config` command-line argument, # where MCP credentials can be observed by process-inspection tools. Passing @@ -337,19 +299,16 @@ def _stage_mcp_config(payload: dict[str, Any]) -> ClaudeMcpSettings | None: # staged file if the process is terminated before runtime cleanup can remove # it. The file remains owner-only as defense in depth and is retained until # the SDK client disconnects normally. - servers = _mcp_servers(payload) + servers = _mcp_servers(config) if not servers: return None - fabric_runtime_id = runtime_id(payload) + fabric_runtime_id = runtime_context.runtime_id environment: dict[str, str] = {} for server_name, server in servers.items(): raw_environment = server.get("env") if raw_environment is None: continue - server_environment = _mapping( - raw_environment, - name=f"MCP server {server_name} env", - ) + server_environment = raw_environment projected_environment: dict[str, str] = {} for variable_name, value in sorted(server_environment.items()): if not isinstance(variable_name, str) or not variable_name: @@ -371,7 +330,7 @@ def _stage_mcp_config(payload: dict[str, Any]) -> ClaudeMcpSettings | None: server["env"] = projected_environment config_root = ( - _artifact_root(payload) + _artifact_root(runtime_context, base_dir) / ".fabric" / "claude" / "mcp" @@ -408,25 +367,15 @@ def _cleanup_mcp_config(config_path: Path | None) -> None: LOGGER.exception("Claude MCP runtime configuration could not be removed") -def _native_skill_paths(payload: dict[str, Any]) -> list[Path]: - native = ( - _mapping(common_utils.capability_plan(payload), name="capability_plan").get( - "native" - ) - or {} - ) - values = _mapping(native, name="capability_plan.native").get("skill_paths") or [] - if not isinstance(values, list) or any( - not isinstance(value, (str, Path)) for value in values - ): - raise AdapterConfigError( - "claude_invalid_configuration", "native skill_paths must be a list of paths" - ) - return [_resolve_path(payload, value) for value in values] +def _native_skill_paths(config: AgentConfig, base_dir: str) -> list[Path]: + values = config.skills.paths if config.skills is not None else [] + return [_resolve_path(base_dir, value) for value in values] -def _stage_skill_plugin(payload: dict[str, Any]) -> list[dict[str, str]]: - skill_paths = _native_skill_paths(payload) +def _stage_skill_plugin( + config: AgentConfig, runtime_context: RuntimeContext, base_dir: str +) -> list[dict[str, str]]: + skill_paths = _native_skill_paths(config, base_dir) if not skill_paths: return [] @@ -447,9 +396,13 @@ def _stage_skill_plugin(payload: dict[str, Any]) -> list[dict[str, str]]: names.add(name) skills.append((name, skill_path)) - plugin_key = sha256(runtime_id(payload).encode()).hexdigest() + plugin_key = sha256(runtime_context.runtime_id.encode()).hexdigest() plugin_root = ( - _artifact_root(payload) / ".fabric" / "claude" / "plugins" / plugin_key + _artifact_root(runtime_context, base_dir) + / ".fabric" + / "claude" + / "plugins" + / plugin_key ) if plugin_root.exists(): shutil.rmtree(plugin_root) @@ -502,15 +455,20 @@ def _stage_relay_plugin(plugin_path: Path, executable: Path) -> None: ) -def prepare_claude_relay(payload: dict[str, Any]) -> ClaudeRelaySettings | None: +def prepare_claude_relay( + payload: dict[str, Any], + model: AgentModelConfig, + runtime_context: RuntimeContext, + base_dir: str, +) -> ClaudeRelaySettings | None: """Generate Relay gateway and Claude hook configuration.""" - if not common_utils.relay_enabled(payload): + if runtime_context.telemetry is None or not runtime_context.telemetry.relay_enabled: return None command = os.environ.get("FABRIC_TEST_NEMO_RELAY_COMMAND", "nemo-relay") try: executable = relay_gateway.resolve_relay_command( - Path(common_utils.base_dir(payload)).resolve(), + Path(base_dir).resolve(), command, ) except FileNotFoundError as error: @@ -521,7 +479,9 @@ def prepare_claude_relay(payload: dict[str, Any]) -> ClaudeRelaySettings | None: try: relay_contract = relay_gateway.relay_cli_contract(executable) - plugin_config = common_utils.load_relay_plugin_config(payload) + plugin_config = common_utils.load_relay_plugin_config( + payload, model_name=model.model + ) config_path, plugin_config_path = common_utils.write_relay_configs( relay_config={"agents": {"claude": {"command": "claude"}}}, plugin_config=plugin_config, @@ -551,7 +511,7 @@ def prepare_claude_relay(payload: dict[str, Any]) -> ClaudeRelaySettings | None: bind=gateway_bind, url=f"http://{gateway_bind}", log_path=config_path.parent / "gateway.log", - anthropic_base_url=_anthropic_base_url(_selected_model_config(payload)), + anthropic_base_url=_anthropic_base_url(model), ) plugin_path = config_path.parent / "claude-plugin" try: @@ -573,11 +533,11 @@ def discard_stderr(_: str) -> None: """Consume Claude Code stderr without exposing it through Fabric artifacts.""" -def tool_policy_hooks(payload: dict[str, Any]) -> dict[str, list[HookMatcher]] | None: +def tool_policy_hooks(config: AgentConfig) -> dict[str, list[HookMatcher]] | None: """Enforce the normalized tool policy across built-in, MCP, and plugin tools.""" - enabled = common_utils.enabled_tools(payload) - blocked = set(common_utils.blocked_tools(payload)) + enabled = config.tools.enabled if config.tools is not None else None + blocked = set(config.tools.blocked if config.tools is not None else []) if enabled is None and not blocked: return None enabled_set = None if enabled is None else set(enabled) @@ -607,17 +567,20 @@ async def enforce_policy( def build_options( - payload: dict[str, Any], + config: AgentConfig, + runtime_context: RuntimeContext, + base_dir: str, *, relay: ClaudeRelaySettings | None = None, ) -> ClaudeAgentOptions: - settings = _settings(payload) + settings = _settings(config) + model = _selected_model_config(config) permission_mode = settings.get("permission_mode") if permission_mode is not None and permission_mode not in PERMISSION_MODES: raise AdapterConfigError( "claude_invalid_configuration", "permission_mode is invalid" ) - max_turns = common_utils.max_turns(payload) + max_turns = config.runtime.max_turns if config.runtime is not None else None max_budget = settings.get("max_budget_usd") if max_budget is not None: max_budget = _positive_number(max_budget, name="max_budget_usd") @@ -629,39 +592,43 @@ def build_options( ) cli_path = os.environ.get("FABRIC_TEST_CLAUDE_CLI_PATH") - system_prompt = common_utils.system_instruction(payload) - enabled_tools = common_utils.enabled_tools(payload) + instructions = config.instructions + system_prompt = ( + instructions.system.content if instructions and instructions.system else None + ) + enabled_tools = config.tools.enabled if config.tools is not None else None allowed_tools = ( enabled_tools if permission_mode == "dontAsk" and enabled_tools is not None else [] ) - plugins = _stage_skill_plugin(payload) + plugins = _stage_skill_plugin(config, runtime_context, base_dir) has_skill_plugin = bool(plugins) if relay is not None: plugins.append({"type": "local", "path": str(relay.plugin_path)}) environment = child_environment( - payload, + model, + runtime_context, relay_gateway_url=relay.gateway.url if relay is not None else None, ) - mcp = _stage_mcp_config(payload) + mcp = _stage_mcp_config(config, runtime_context, base_dir) if mcp is not None: environment.update(mcp.environment) try: return ClaudeAgentOptions( - cwd=resolve_cwd(payload), - model=selected_model(payload), + cwd=resolve_cwd(runtime_context, base_dir), + model=selected_model(model), system_prompt=system_prompt, tools=enabled_tools, allowed_tools=allowed_tools, - disallowed_tools=common_utils.blocked_tools(payload), - hooks=tool_policy_hooks(payload), + disallowed_tools=config.tools.blocked if config.tools is not None else [], + hooks=tool_policy_hooks(config), permission_mode=permission_mode, max_turns=max_turns, max_budget_usd=max_budget, setting_sources=sources, - cli_path=_resolve_path(payload, cli_path) if cli_path else None, + cli_path=_resolve_path(base_dir, cli_path) if cli_path else None, mcp_servers=mcp.config_path if mcp is not None else {}, strict_mcp_config=True, skills="all" if has_skill_plugin else None, @@ -674,17 +641,15 @@ def build_options( raise -def timeout_seconds(payload: dict[str, Any]) -> float: - value = common_utils.timeout_seconds(payload, default=1800) - return _positive_number(value, name="timeout_seconds") +def timeout_seconds() -> float: + return 1800.0 -def _artifact_root(payload: dict[str, Any]) -> Path: - artifacts = common_utils.runtime_context(payload).get("artifacts") or {} - root = artifacts.get("root") if isinstance(artifacts, dict) else None +def _artifact_root(runtime_context: RuntimeContext, base_dir: str) -> Path: + root = runtime_context.artifacts.root if root: return Path(root) - return Path(common_utils.base_dir(payload)) / "artifacts" / "claude" + return Path(base_dir) / "artifacts" / "claude" def _json_safe(value: Any) -> Any: @@ -714,9 +679,8 @@ def _result_failed(result: ResultMessage) -> bool: def normalize_result( - payload: dict[str, Any], messages: list[Message], result: ResultMessage + messages: list[Message], result: ResultMessage ) -> dict[str, Any]: - del payload failed = _result_failed(result) error = None if failed: @@ -786,7 +750,8 @@ def sdk_failure(error: BaseException) -> dict[str, Any]: def child_environment( - payload: dict[str, Any], + model: AgentModelConfig, + runtime_context: RuntimeContext, *, relay_gateway_url: str | None = None, ) -> dict[str, str]: @@ -794,13 +759,12 @@ def child_environment( values.update( {name: os.environ[name] for name in INHERITED_ENV_NAMES if name in os.environ} ) - model = _selected_model_config(payload) - api_key_env = model.get("api_key_env") - if isinstance(api_key_env, str) and api_key_env in os.environ: + api_key_env = model.api_key_env + if api_key_env is not None and api_key_env in os.environ: values[api_key_env] = os.environ[api_key_env] - configured = common_utils.environment_env(payload) + configured = runtime_context.environment.env values.update(configured) - model_environment = _model_environment(payload, values) + model_environment = _model_environment(model, values) conflicts = sorted( name for name, value in model_environment.items() @@ -844,7 +808,8 @@ def _relay_output( def _start_relay_gateway( - payload: dict[str, Any], + runtime_context: RuntimeContext, + base_dir: str, relay: ClaudeRelaySettings | None, ) -> subprocess.Popen[Any] | None: if relay is None: @@ -852,7 +817,7 @@ def _start_relay_gateway( try: return relay_gateway.start_relay_gateway( launch=relay.gateway, - cwd=resolve_cwd(payload), + cwd=resolve_cwd(runtime_context, base_dir), ) except relay_gateway.RelayGatewayError as error: raise AdapterRelayError( @@ -926,7 +891,7 @@ class ClaudeRuntime: """One connected Claude SDK client owned by a Fabric runtime.""" def __init__(self) -> None: - self._start_payload: dict[str, Any] | None = None + self._agent_config: AgentConfig | None = None self._fabric_runtime_id: str | None = None self._claude_session_id: str | None = None self._client: ClaudeSDKClient | None = None @@ -942,11 +907,19 @@ async def start(self, payload: dict[str, Any]) -> None: "Claude runtime is already started", ) try: - fabric_runtime_id = runtime_id(payload) - relay = prepare_claude_relay(payload) + agent_config = payload.get("config") + if not isinstance(agent_config, AgentConfig): + raise lifecycle.LifecycleError( + "claude_invalid_config", "Claude requires a validated AgentConfig" + ) + runtime_context = _runtime_context(payload) + base_dir = common_utils.base_dir(payload) + model = _selected_model_config(agent_config) + fabric_runtime_id = runtime_context.runtime_id + relay = prepare_claude_relay(payload, model, runtime_context, base_dir) self._relay = relay - self._gateway_process = _start_relay_gateway(payload, relay) - options = build_options(payload, relay=relay) + self._gateway_process = _start_relay_gateway(runtime_context, base_dir, relay) + options = build_options(agent_config, runtime_context, base_dir, relay=relay) if isinstance(options.mcp_servers, Path): self._mcp_config_path = options.mcp_servers client = ClaudeSDKClient(options) @@ -961,29 +934,25 @@ async def start(self, payload: dict[str, Any]) -> None: self._cleanup_failed_start() raise - self._start_payload = payload + self._agent_config = agent_config self._fabric_runtime_id = fabric_runtime_id self._client = client async def invoke(self, invocation: dict[str, Any]) -> dict[str, Any]: client = self._client - start_payload = self._start_payload + agent_config = self._agent_config fabric_runtime_id = self._fabric_runtime_id - if client is None or start_payload is None or fabric_runtime_id is None: + if client is None or agent_config is None or fabric_runtime_id is None: raise lifecycle.LifecycleError( "claude_runtime_not_started", "Claude runtime is not started", ) - if runtime_id(invocation) != fabric_runtime_id: + runtime_context = _runtime_context(invocation) + if runtime_context.runtime_id != fabric_runtime_id: raise lifecycle.LifecycleError( "claude_runtime_mismatch", "Claude invocation does not match the connected runtime", ) - payload = { - **start_payload, - "runtime_context": invocation.get("runtime_context"), - "request": invocation.get("request"), - } if self._unusable: return _failure( "claude_runtime_unavailable", @@ -991,8 +960,8 @@ async def invoke(self, invocation: dict[str, Any]) -> dict[str, Any]: ) try: - prompt = request_prompt(payload) - invocation_timeout = timeout_seconds(payload) + prompt = request_prompt(invocation) + invocation_timeout = timeout_seconds() except ClaudeAdapterError as error: output = adapter_failure(error) else: @@ -1004,7 +973,6 @@ async def invoke(self, invocation: dict[str, Any]) -> dict[str, Any]: else None ) output = await self._run_query( - payload, client, prompt, invocation_timeout, @@ -1039,7 +1007,6 @@ async def invoke(self, invocation: dict[str, Any]) -> dict[str, Any]: async def _run_query( self, - payload: dict[str, Any], client: ClaudeSDKClient, prompt: str, invocation_timeout: float, @@ -1067,7 +1034,7 @@ async def _run_query( if result is None or not _result_failed(result): raise LOGGER.exception("Claude SDK stream raised after a failed terminal result") - output = self._normalize_invocation(payload, messages, result) + output = self._normalize_invocation(messages, result) else: if result is None: self._unusable = True @@ -1075,17 +1042,16 @@ async def _run_query( "claude_missing_result", "Claude returned no terminal result" ) else: - output = self._normalize_invocation(payload, messages, result) + output = self._normalize_invocation(messages, result) return output def _normalize_invocation( self, - payload: dict[str, Any], messages: list[Message], result: ResultMessage, ) -> dict[str, Any]: try: - output = normalize_result(payload, messages, result) + output = normalize_result(messages, result) if not output["failed"]: invalid_session = _validate_result_session( self._claude_session_id, result @@ -1102,7 +1068,7 @@ def _normalize_invocation( async def stop(self) -> None: client = self._client self._client = None - self._start_payload = None + self._agent_config = None self._fabric_runtime_id = None self._claude_session_id = None self._unusable = True @@ -1161,7 +1127,7 @@ async def _interrupt_failed_invocation(self) -> None: def main() -> None: """Serve the persistent local-host lifecycle protocol.""" - lifecycle.serve(ClaudeRuntime) + lifecycle.serve(ClaudeRuntime, config_loader=AgentConfig.from_mapping) if __name__ == "__main__": diff --git a/adapters/claude/uv.lock b/adapters/claude/uv.lock index 22798240..47f64c25 100644 --- a/adapters/claude/uv.lock +++ b/adapters/claude/uv.lock @@ -348,11 +348,21 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/01/c8/248b201f6d753d69fd5d6506011abbb35a946d9142b2ae311a948fd0be3d/mcp-1.29.0-py3-none-any.whl", hash = "sha256:f5a075bb611f23d6f4d080c6a1699fa62772eebc562ba9e66b306ddde1c755f7", size = 223436, upload-time = "2026-07-28T13:41:40.337Z" }, ] +[[package]] +name = "nemo-fabric-adapter-contract" +version = "0.2.0" +source = { editable = "../../adapter-contract" } + +[package.metadata] +requires-dist = [{ name = "pydantic", marker = "extra == 'pydantic'", specifier = ">=2.12,<3" }] +provides-extras = ["pydantic"] + [[package]] name = "nemo-fabric-adapters-claude" version = "0.2.0" source = { editable = "." } dependencies = [ + { name = "nemo-fabric-adapter-contract" }, { name = "nemo-fabric-adapters-common" }, { name = "tomli-w" }, ] @@ -369,6 +379,7 @@ harness = [ requires-dist = [ { name = "claude-agent-sdk", marker = "extra == 'full'", specifier = "==0.2.120" }, { name = "claude-agent-sdk", marker = "extra == 'harness'", specifier = "==0.2.120" }, + { name = "nemo-fabric-adapter-contract", editable = "../../adapter-contract" }, { name = "nemo-fabric-adapters-common", editable = "../common" }, { name = "tomli-w", specifier = "~=1.2" }, ] diff --git a/adapters/common/src/nemo_fabric_adapters/common/utils.py b/adapters/common/src/nemo_fabric_adapters/common/utils.py index e44d1cc6..3f85d7d2 100644 --- a/adapters/common/src/nemo_fabric_adapters/common/utils.py +++ b/adapters/common/src/nemo_fabric_adapters/common/utils.py @@ -233,7 +233,9 @@ def dump_yaml(value: dict[str, Any]) -> str: return json.dumps(value, indent=2, sort_keys=False) + "\n" -def load_relay_plugin_config(payload: dict[str, Any]) -> dict[str, Any]: +def load_relay_plugin_config( + payload: dict[str, Any], *, model_name: str | None = None +) -> dict[str, Any]: config_path = os.environ.get("FABRIC_RELAY_CONFIG_PATH") if not config_path: raise RuntimeError("FABRIC_RELAY_CONFIG_PATH is required when Relay is enabled") @@ -256,12 +258,12 @@ def load_relay_plugin_config(payload: dict[str, Any]) -> dict[str, Any]: } plugin_config.setdefault("version", 1) plugin_config.setdefault("components", []) - normalize_relay_output_dirs(plugin_config, payload) + normalize_relay_output_dirs(plugin_config, payload, model_name=model_name) return plugin_config def normalize_relay_output_dirs( - plugin_config: dict[str, Any], payload: dict[str, Any] + plugin_config: dict[str, Any], payload: dict[str, Any], *, model_name: str | None = None ) -> None: base = Path(base_dir(payload)).resolve() runtime_id = runtime_context(payload)["runtime_id"] @@ -303,7 +305,7 @@ def normalize_relay_output_dirs( Path(atif["output_directory"]).mkdir(parents=True, exist_ok=True) atif.setdefault("filename_template", "trajectory-{session_id}.atif.json") atif.setdefault("agent_name", agent_name(payload)) - atif.setdefault("model_name", relay_model_name(payload)) + atif.setdefault("model_name", model_name or relay_model_name(payload)) def _artifact_directory(value: Any) -> Path | None: diff --git a/crates/fabric-cli/assets/adapters/claude/fabric-adapter.json b/crates/fabric-cli/assets/adapters/claude/fabric-adapter.json index d04516c6..56710f05 100644 --- a/crates/fabric-cli/assets/adapters/claude/fabric-adapter.json +++ b/crates/fabric-cli/assets/adapters/claude/fabric-adapter.json @@ -41,6 +41,7 @@ "additionalProperties": false }, "config": { + "input": "agent_config", "accepts": [ "models", "models.base_url", diff --git a/examples/harbor/swebench/adapters/claude/fabric-adapter.json b/examples/harbor/swebench/adapters/claude/fabric-adapter.json index d04516c6..56710f05 100644 --- a/examples/harbor/swebench/adapters/claude/fabric-adapter.json +++ b/examples/harbor/swebench/adapters/claude/fabric-adapter.json @@ -41,6 +41,7 @@ "additionalProperties": false }, "config": { + "input": "agent_config", "accepts": [ "models", "models.base_url", diff --git a/tests/adapters/test_adapter_package_metadata.py b/tests/adapters/test_adapter_package_metadata.py index 8ed912fa..5f4066f5 100644 --- a/tests/adapters/test_adapter_package_metadata.py +++ b/tests/adapters/test_adapter_package_metadata.py @@ -64,6 +64,7 @@ def load_pyproject(path: str) -> dict: ( "adapters/claude", [ + f"nemo-fabric-adapter-contract == {PACKAGE_VERSION}", f"nemo-fabric-adapters-common == {PACKAGE_VERSION}", "tomli-w~=1.2", ], diff --git a/tests/adapters/test_claude_adapter.py b/tests/adapters/test_claude_adapter.py index 4ecbdc8d..759941e2 100644 --- a/tests/adapters/test_claude_adapter.py +++ b/tests/adapters/test_claude_adapter.py @@ -27,6 +27,8 @@ from claude_agent_sdk import SystemMessage from claude_agent_sdk import TextBlock from claude_agent_sdk._errors import MessageParseError +from nemo_fabric_adapter_contract.models import AgentConfig +from nemo_fabric_adapter_contract.models import RuntimeContext from nemo_fabric_adapters.claude import adapter ROOT = Path(__file__).resolve().parents[2] @@ -46,11 +48,52 @@ def lifecycle_invocation(payload: dict[str, Any]) -> dict[str, Any]: return { - "runtime_context": payload["runtime_context"], + "runtime_context": { + **payload["runtime_context"], + "request_id": payload["request"]["request_id"], + }, "request": payload["request"], } +def agent_config(payload: dict[str, Any]) -> AgentConfig: + return AgentConfig.from_mapping(payload["config"]) + + +def runtime_context(payload: dict[str, Any]) -> RuntimeContext: + return RuntimeContext.from_mapping( + {**payload["runtime_context"], "request_id": payload["request"]["request_id"]} + ) + + +def build_options(payload: dict[str, Any], *, relay=None) -> ClaudeAgentOptions: + return adapter.build_options( + agent_config(payload), runtime_context(payload), payload["base_dir"], relay=relay + ) + + +def prepare_claude_relay(payload: dict[str, Any]): + config = agent_config(payload) + return adapter.prepare_claude_relay( + payload, + adapter._selected_model_config(config), + runtime_context(payload), + payload["base_dir"], + ) + + +def lifecycle_start(payload: dict[str, Any]) -> dict[str, Any]: + return { + **payload, + "config": agent_config(payload), + "runtime_context": { + **payload["runtime_context"], + "request_id": (payload.get("request") or {}).get("request_id", "request-1"), + }, + "request": None, + } + + def install_fake_client( monkeypatch: pytest.MonkeyPatch, response_factory: Callable[[MagicMock], AsyncIterator[Message]], @@ -124,6 +167,7 @@ def test_claude_descriptor_is_narrow_and_versioned(): "additionalProperties": False, }, "config": { + "input": "agent_config", "accepts": [ "models", "models.base_url", @@ -158,7 +202,6 @@ def claude_payload_fixture(tmp_path) -> dict[str, Any]: "base_dir": str(tmp_path), "config": { "harness": { - "adapter_id": "nvidia.fabric.claude", "settings": { "permission_mode": "dontAsk", "max_budget_usd": 1.5, @@ -168,7 +211,7 @@ def claude_payload_fixture(tmp_path) -> dict[str, Any]: "instructions": { "system": {"content": "Review carefully.", "mode": "replace"} }, - "runtime": {"timeout_seconds": 30, "max_turns": 4}, + "runtime": {"max_turns": 4}, "models": { "default": { "provider": "anthropic", @@ -177,11 +220,30 @@ def claude_payload_fixture(tmp_path) -> dict[str, Any]: } }, "tools": {"blocked": ["Bash"]}, + "skills": {"paths": [str(skill_path)]}, + "mcp": { + "servers": { + "repo": { + "transport": "stdio", + "url": "repo-mcp", + "args": ["--root", ".", "--config", "repo config.json"], + "env": {"REPO_MCP_MODE": "mcp-secret-value"}, + }, + "docs": { + "transport": "streamable-http", + "url": "https://mcp.example.test", + }, + } + }, }, "runtime_context": { "runtime_id": "runtime-claude-1", "invocation_id": "invocation-1", "environment": { + "environment_id": "environment-claude-1", + "provider": "local", + "control_location": "in_env_control", + "ownership": "caller_owned", "workspace": str(workspace), "env": {"ANTHROPIC_API_KEY": "configured-secret"}, }, @@ -212,7 +274,7 @@ def claude_payload_fixture(tmp_path) -> dict[str, Any]: def test_build_options_maps_normalized_capabilities_and_claude_settings(claude_payload): - options = adapter.build_options(claude_payload) + options = build_options(claude_payload) assert options.cwd == Path( claude_payload["runtime_context"]["environment"]["workspace"] ) @@ -270,7 +332,7 @@ async def test_tool_policy_hooks_gate_built_in_and_mcp_tools(claude_payload): "blocked": ["Bash"], } - options = adapter.build_options(claude_payload) + options = build_options(claude_payload) assert options.tools == ["Read", "Edit"] assert options.allowed_tools == ["Read", "Edit"] @@ -291,7 +353,7 @@ def test_enabled_tools_do_not_populate_allowed_tools_in_default_mode(claude_payl claude_payload["config"]["harness"]["settings"]["permission_mode"] = "default" claude_payload["config"]["tools"] = {"enabled": ["Read"], "blocked": []} - options = adapter.build_options(claude_payload) + options = build_options(claude_payload) assert options.permission_mode == "default" assert options.tools == ["Read"] @@ -315,9 +377,9 @@ def relay_payload_fixture(claude_payload, tmp_path) -> dict[str, Any]: encoding="utf-8", ) os.environ["FABRIC_RELAY_CONFIG_PATH"] = str(relay_intent_path) - claude_payload["telemetry_plan"] = { - "providers": ["relay"], + claude_payload["runtime_context"]["telemetry"] = { "relay_enabled": True, + "metadata": {"telemetry_providers": ["relay"]}, } return claude_payload @@ -354,7 +416,7 @@ def test_prepare_claude_relay_writes_gateway_config_and_complete_hook_plugin( ), ) - relay = adapter.prepare_claude_relay(relay_payload) + relay = prepare_claude_relay(relay_payload) assert relay is not None assert relay.gateway.executable == executable @@ -429,9 +491,9 @@ def test_build_options_adds_relay_plugin_and_gateway_environment( ) ), ) - relay = adapter.prepare_claude_relay(relay_payload) + relay = prepare_claude_relay(relay_payload) - options = adapter.build_options(relay_payload, relay=relay) + options = build_options(relay_payload, relay=relay) assert options.env["NEMO_RELAY_GATEWAY_URL"] == relay.gateway.url assert options.env["ANTHROPIC_BASE_URL"] == relay.gateway.url @@ -445,7 +507,7 @@ def test_build_options_adds_relay_plugin_and_gateway_environment( def test_build_options_does_not_enable_skills_for_relay_plugin_alone( relay_payload, tmp_path ): - relay_payload["capability_plan"]["native"]["skill_paths"] = [] + relay_payload["config"]["skills"]["paths"] = [] relay = adapter.ClaudeRelaySettings( gateway=adapter.relay_gateway.RelayGatewayLaunch( executable=tmp_path / "nemo-relay", @@ -458,7 +520,7 @@ def test_build_options_does_not_enable_skills_for_relay_plugin_alone( plugin_path=tmp_path / "relay-plugin", ) - options = adapter.build_options(relay_payload, relay=relay) + options = build_options(relay_payload, relay=relay) assert options.tools is None assert options.skills is None @@ -468,18 +530,25 @@ def test_build_options_does_not_enable_skills_for_relay_plugin_alone( def test_build_options_maps_blocked_tools_to_disallowed_tools(claude_payload): claude_payload["config"]["tools"] = {"blocked": ["Bash", "WebFetch"]} - options = adapter.build_options(claude_payload) + options = build_options(claude_payload) assert options.tools is None assert options.disallowed_tools == ["Bash", "WebFetch"] def test_build_options_rejects_skill_path_without_skill_manifest(claude_payload): - skill_path = Path(claude_payload["capability_plan"]["native"]["skill_paths"][0]) + skill_path = Path(claude_payload["config"]["skills"]["paths"][0]) (skill_path / "SKILL.md").unlink() - with pytest.raises(adapter.AdapterConfigError, match="SKILL.md"): - adapter.build_options(claude_payload) + with pytest.raises(adapter.AdapterConfigError, match=r"SKILL\.md"): + build_options(claude_payload) + + +def test_runtime_context_validation_uses_lifecycle_error(): + with pytest.raises(adapter.lifecycle.LifecycleError) as caught: + adapter._runtime_context({"runtime_context": {}}) + + assert caught.value.code == "claude_invalid_runtime_context" def test_build_options_maps_custom_provider_to_claude_gateway_environment( @@ -497,7 +566,7 @@ def test_build_options_maps_custom_provider_to_claude_gateway_environment( claude_payload["runtime_context"]["environment"]["env"].pop("ANTHROPIC_API_KEY") os.environ["ACME_API_KEY"] = "acme-secret" - options = adapter.build_options(claude_payload) + options = build_options(claude_payload) assert options.model == "aws/anthropic/claude-opus-4-5" assert options.env["ANTHROPIC_BASE_URL"] == "https://acme.example" @@ -520,7 +589,7 @@ def test_build_options_requires_custom_provider_api_key_env( claude_payload["runtime_context"]["environment"]["env"].pop("ANTHROPIC_API_KEY") with pytest.raises(adapter.AdapterConfigError, match="api_key_env is required"): - adapter.build_options(claude_payload) + build_options(claude_payload) def test_build_options_requires_custom_provider_credential(claude_payload): @@ -536,7 +605,7 @@ def test_build_options_requires_custom_provider_credential(claude_payload): os.environ.pop("ACME_API_KEY", None) with pytest.raises(adapter.AdapterConfigError, match="ACME_API_KEY is required"): - adapter.build_options(claude_payload) + build_options(claude_payload) def test_build_options_requires_custom_provider_endpoint(claude_payload): @@ -553,7 +622,7 @@ def test_build_options_requires_custom_provider_endpoint(claude_payload): } with pytest.raises(adapter.AdapterConfigError, match="base_url is required"): - adapter.build_options(claude_payload) + build_options(claude_payload) @pytest.mark.parametrize( @@ -584,15 +653,15 @@ def test_build_options_rejects_model_environment_conflicts( adapter.AdapterConfigError, match=rf"environment\.env\.{name} conflicts", ): - adapter.build_options(claude_payload) + build_options(claude_payload) def test_selected_model_rejects_empty_provider(claude_payload): model = claude_payload["config"]["models"]["default"] model["provider"] = "" - with pytest.raises(adapter.AdapterConfigError, match="non-empty string"): - adapter.selected_model(claude_payload) + with pytest.raises(Exception, match="non-empty lowercase identifier"): + adapter.selected_model(adapter._selected_model_config(agent_config(claude_payload))) def test_normalize_result_exposes_session_usage_cost_and_buffered_events( @@ -619,7 +688,7 @@ def test_normalize_result_exposes_session_usage_cost_and_buffered_events( result="done", ) - output = adapter.normalize_result(claude_payload, messages, result) + output = adapter.normalize_result(messages, result) assert output["response"] == "done" assert output["session_id"] == "claude-session" @@ -678,7 +747,7 @@ async def interrupt(self): start_payload = dict(claude_payload) start_payload.pop("request") runtime = adapter.ClaudeRuntime() - await runtime.start(start_payload) + await runtime.start(lifecycle_start(start_payload)) mcp_config_path = clients[0].options.mcp_servers assert isinstance(mcp_config_path, Path) assert mcp_config_path.exists() @@ -721,7 +790,7 @@ async def connect(self): start_payload.pop("request") with pytest.raises(adapter.lifecycle.LifecycleError) as caught: - await adapter.ClaudeRuntime().start(start_payload) + await adapter.ClaudeRuntime().start(lifecycle_start(start_payload)) assert caught.value.code == "claude_connection_failed" assert len(staged_paths) == 1 @@ -784,7 +853,7 @@ async def interrupt(self): start_payload = dict(relay_payload) start_payload.pop("request") runtime = adapter.ClaudeRuntime() - await runtime.start(start_payload) + await runtime.start(lifecycle_start(start_payload)) first = await runtime.invoke(lifecycle_invocation(relay_payload)) relay_payload["runtime_context"]["invocation_id"] = "invocation-2" second = await runtime.invoke(lifecycle_invocation(relay_payload)) @@ -896,7 +965,7 @@ async def write_atif(): start_payload = { key: value for key, value in relay_payload.items() if key != "request" } - await runtime.start(start_payload) + await runtime.start(lifecycle_start(start_payload)) try: output = await runtime.invoke(lifecycle_invocation(relay_payload)) assert write_task is not None @@ -953,7 +1022,7 @@ async def responses(_client) -> AsyncIterator[ResultMessage]: start_payload = { key: value for key, value in relay_payload.items() if key != "request" } - await runtime.start(start_payload) + await runtime.start(lifecycle_start(start_payload)) try: output = await runtime.invoke(lifecycle_invocation(relay_payload)) unavailable = await runtime.invoke(lifecycle_invocation(relay_payload)) @@ -1024,7 +1093,7 @@ async def responses(_client) -> AsyncIterator[ResultMessage]: start_payload = { key: value for key, value in relay_payload.items() if key != "request" } - await runtime.start(start_payload) + await runtime.start(lifecycle_start(start_payload)) output = await runtime.invoke(lifecycle_invocation(relay_payload)) with pytest.raises(adapter.lifecycle.LifecycleError) as caught: await runtime.stop() @@ -1082,7 +1151,7 @@ async def responses(_client) -> AsyncIterator[ResultMessage]: start_payload = { key: value for key, value in relay_payload.items() if key != "request" } - await runtime.start(start_payload) + await runtime.start(lifecycle_start(start_payload)) output = await runtime.invoke(lifecycle_invocation(relay_payload)) with pytest.raises(adapter.lifecycle.LifecycleError) as caught: await runtime.stop() @@ -1133,7 +1202,7 @@ async def responses(_client) -> AsyncIterator[ResultMessage]: start_payload = { key: value for key, value in relay_payload.items() if key != "request" } - await runtime.start(start_payload) + await runtime.start(lifecycle_start(start_payload)) try: if isinstance(failure, asyncio.CancelledError): with pytest.raises(asyncio.CancelledError): @@ -1177,7 +1246,7 @@ async def responses(_client) -> AsyncIterator[ResultMessage]: start_payload = { key: value for key, value in claude_payload.items() if key != "request" } - await runtime.start(start_payload) + await runtime.start(lifecycle_start(start_payload)) output = await runtime.invoke(lifecycle_invocation(claude_payload)) await runtime.stop() @@ -1225,7 +1294,7 @@ async def test_runtime_start_reports_relay_failure_without_raw_diagnostic( key: value for key, value in relay_payload.items() if key != "request" } with pytest.raises(adapter.lifecycle.LifecycleError) as caught: - await runtime.start(start_payload) + await runtime.start(lifecycle_start(start_payload)) assert caught.value.code == "claude_relay_start_failed" assert caught.value.message == "NeMo Relay gateway failed to start" @@ -1278,7 +1347,7 @@ def test_build_options_forwards_anthropic_auth_environment( os.environ["FABRIC_UNRELATED_SECRET"] = "do-not-forward" os.environ.update(auth_environment) - options = adapter.build_options(claude_payload) + options = build_options(claude_payload) forwarded_auth_environment = { name: options.env[name] @@ -1294,7 +1363,7 @@ def test_build_options_preserves_unix_user_for_cached_login( ): os.environ["USER"] = "fabric-user" - options = adapter.build_options(claude_payload) + options = build_options(claude_payload) assert options.env["USER"] == "fabric-user" @@ -1336,7 +1405,7 @@ def test_error_result_is_normalized_as_failure(claude_payload): errors=["provider-specific failure"], ) - output = adapter.normalize_result(claude_payload, [], result) + output = adapter.normalize_result([], result) assert output["failed"] is True assert output["error"] == { @@ -1357,7 +1426,7 @@ def test_error_subtype_is_failure_when_sdk_flag_is_false(claude_payload): session_id="claude-session", ) - output = adapter.normalize_result(claude_payload, [], result) + output = adapter.normalize_result([], result) assert output["completed"] is False assert output["failed"] is True @@ -1370,4 +1439,4 @@ def test_main_serves_persistent_runtime(monkeypatch): adapter.main() - serve.assert_called_once_with(adapter.ClaudeRuntime) + serve.assert_called_once_with(adapter.ClaudeRuntime, config_loader=AgentConfig.from_mapping) diff --git a/uv.lock b/uv.lock index 2c10184e..cf0f9154 100644 --- a/uv.lock +++ b/uv.lock @@ -2255,6 +2255,7 @@ name = "nemo-fabric-adapters-claude" version = "0.2.0" source = { editable = "adapters/claude" } dependencies = [ + { name = "nemo-fabric-adapter-contract" }, { name = "nemo-fabric-adapters-common" }, { name = "tomli-w" }, ] @@ -2268,6 +2269,7 @@ harness = [ requires-dist = [ { name = "claude-agent-sdk", marker = "extra == 'full'", specifier = "==0.2.120" }, { name = "claude-agent-sdk", marker = "extra == 'harness'", specifier = "==0.2.120" }, + { name = "nemo-fabric-adapter-contract", editable = "adapter-contract" }, { name = "nemo-fabric-adapters-common", editable = "adapters/common" }, { name = "tomli-w", specifier = "~=1.2" }, ]