diff --git a/docs/mcp-tools.md b/docs/mcp-tools.md index 1b2f40af..b9e93a59 100644 --- a/docs/mcp-tools.md +++ b/docs/mcp-tools.md @@ -215,6 +215,52 @@ uv run conductor run workflow.yaml --input.question="test" > ``` > See [copilot-sdk#163](https://github.com/github/copilot-sdk/issues/163) for status. +## Working Directory + +You can specify a default directory where stdio-based MCP servers and agent sessions execute. This is useful for workflows interacting with local repositories or specific folders. + +The working directory is configured in the workflow `runtime` block or on individual agents: + +```yaml +workflow: + runtime: + working_dir: "/path/to/default/workspace" # Global default +``` + +Or on a specific agent: + +```yaml +agents: + - name: code_expert + working_dir: "/path/to/specific/repo" +``` + +### Precedence and Resolution + +When determining the active working directory, Conductor follows this precedence: + +1. **Agent level:** The agent's own `working_dir` configuration. +2. **Runtime level:** The global `workflow.runtime.working_dir` default. +3. **Fallback:** If neither is set, Conductor falls back to the current directory of the parent process (`os.getcwd()`). + +Both levels support dynamic values using Jinja2 templates. You can resolve the path at runtime using outputs from previous steps: + +```yaml +agents: + - name: find_repo + type: set + value: "/repositories/my-project" + + - name: git_agent + working_dir: "{{ find_repo.output }}" + prompt: "List the last commits in the repository." +``` + +Relative paths in `working_dir` resolve against the parent directory of the workflow YAML file. If the workflow has no path (e.g. constructed dynamically in memory), they resolve against the current process directory. Conductor lexically normalizes the resolved path. A missing target directory causes Conductor to raise an execution error before any provider call. + +> ⚠️ **Warning: Working directory is NOT a sandbox** +> Setting the working directory doesn't restrict filesystem access. It only sets the default path where the agent session and stdio MCP subprocesses run. The model can still read or write files outside this directory if it uses absolute paths or parent directory traversals (e.g., `../`). Avoid relying on this setting to sandbox untrusted model execution. + ## OAuth Authentication (HTTP/SSE) For HTTP and SSE servers that require OAuth, Conductor can automatically discover OAuth requirements and fetch Azure AD tokens. diff --git a/docs/workflow-syntax.md b/docs/workflow-syntax.md index eb1d572a..2adf23dc 100644 --- a/docs/workflow-syntax.md +++ b/docs/workflow-syntax.md @@ -69,6 +69,10 @@ workflow: # provider-backed agent unless it declares # its own `context_tier`. # See docs/configuration.md#context-tier. + + working_dir: "/path/to/cwd" # Optional: global default working directory for LLM agents + # and their MCP servers. Relative paths resolve against the + # parent directory of the workflow YAML file. ``` **Workflow metadata** is included verbatim in the `workflow_started` event and lets downstream consumers (dashboards, queue runners, observability tools) adapt without parsing the YAML. CLI `--metadata key=value` flags merge on top of YAML metadata (CLI wins on conflicts). @@ -272,6 +276,43 @@ agents: `output_mode` is only valid on provider-backed agents (the default type). It cannot be set on `script`, `human_gate`, or `workflow` agents. +### Working Directory + +Regular LLM agents (provider-backed agents) and their MCP servers run in a specific working directory: + +```yaml +agents: + - name: repository_analyst + working_dir: "./my-project-repo" # Optional: working directory (Jinja2 template) + prompt: | + Examine the repository files and list any issues. +``` + +The `working_dir` field can be defined globally in `workflow.runtime.working_dir` or overridden on individual agents. + +#### Precedence and Path Resolution + +1. **Precedence:** The agent-level `working_dir` overrides the global `workflow.runtime.working_dir`. If neither is configured, the current directory of the parent process (`os.getcwd()`) is used. +2. **Jinja2 Rendering:** Both agent-level and runtime-level configurations support Jinja2 template rendering. This allows dynamic paths, such as directories derived from previous steps: `working_dir: "{{ find_repo.output.path }}"`. +3. **Relative Paths:** Relative paths are resolved against the directory containing the workflow YAML file. When the workflow file location is unknown, relative paths resolve against the current process directory. +4. **Lexical Normalization:** Paths are normalized lexically using `os.path.normpath`. The engine does not resolve symlinks dynamically. + +#### Symlink Semantics + +Because paths are normalized lexically instead of resolving to their real paths: +- Different symlink aliases pointing to the same folder are treated as distinct paths. +- For the Claude provider, distinct paths trigger separate MCP manager connections. This spawns separate MCP server subprocesses for each unique path alias. + +#### Key Restrictions and Exclusions + +- **Rejected Step Types:** The `working_dir` field is strictly rejected on `wait`, `set`, `terminate`, `human_gate`, and `workflow` (sub-workflow) step types. Defining `working_dir` on these steps raises a `ValidationError` at load time. +- **Script Steps:** `script` steps honor only their own `working_dir` field, rendered as a Jinja2 template. `workflow.runtime.working_dir` is not applied; relative paths are passed to the subprocess as-is and therefore resolve against the Conductor process cwd, not the workflow file directory; missing directories surface as subprocess startup `ExecutionError`s rather than the LLM-agent pre-provider working-dir check. +- **Dialog Turns:** The working directory isn't applied to dialog turns in the current version. Multi-turn interactions run in the process default directory. +- **Sub-Workflows:** A sub-workflow doesn't inherit the parent's working directory configuration. Instead, any relative paths in the child workflow resolve against the child workflow's own file directory. + +> ⚠️ **Warning: Working directory is NOT a sandbox** +> Setting `working_dir` doesn't restrict the model's filesystem access. The model can still read and write files outside this directory if it uses absolute paths or parent directory traversals (e.g., `../`). Avoid relying on this configuration to sandbox untrusted model execution. + ### Human Gates Human gates pause workflow execution for user input: diff --git a/examples/working-dir.yaml b/examples/working-dir.yaml new file mode 100644 index 00000000..ca675200 --- /dev/null +++ b/examples/working-dir.yaml @@ -0,0 +1,48 @@ +# Working Directory Example +# +# Demonstrates dynamically setting the working directory for an agent and its +# MCP servers using the outputs of a set step. +# +# The workflow: +# 1. set_step - a set step that calculates the target directory path. +# 2. llm_agent - an LLM agent running with its `working_dir` bound to the +# path output of the set step, using an MCP tool in that directory. +# +# Usage: +# conductor run examples/working-dir.yaml + +workflow: + name: working-dir-demo + description: "Demonstrates dynamic working directory configuration for agents and MCP servers" + version: "1.0.0" + entry_point: set_step + + runtime: + provider: copilot + mcp_servers: + mock-search: + command: npx + args: ["-y", "open-websearch@latest"] + tools: ["*"] + +agents: + - name: set_step + type: set + description: Calculate the directory path for the LLM agent + values: + path: "." # Resolves to the directory containing the workflow YAML (examples/) + routes: + - to: llm_agent + + - name: llm_agent + description: Runs in the dynamically calculated working directory + model: claude-haiku-4.5 + working_dir: "{{ set_step.output.path }}" + prompt: | + You are running in the working directory: {{ set_step.output.path }}. + Please analyze the files in this directory. + routes: + - to: $end + +output: + resolved_path: "{{ set_step.output.path }}" diff --git a/src/conductor/cli/run.py b/src/conductor/cli/run.py index 523258a2..44ff366d 100644 --- a/src/conductor/cli/run.py +++ b/src/conductor/cli/run.py @@ -2220,6 +2220,11 @@ async def resume_workflow_async( # Pass stored session IDs to registry for Copilot session resume if cp.copilot_session_ids: registry.set_resume_session_ids(cp.copilot_session_ids) + # Pass the sessions' original working directories so the provider + # can skip resuming a session whose cwd changed since creation. + # Pre-cwd checkpoints carry an empty mapping (legacy behavior). + if cp.copilot_session_cwds: + registry.set_resume_session_cwds(cp.copilot_session_cwds) # Set up interrupt listener if interactive mode is enabled # Disabled in --web mode since the CLI isn't used for interaction diff --git a/src/conductor/config/schema.py b/src/conductor/config/schema.py index 303efdc6..4cd7260d 100644 --- a/src/conductor/config/schema.py +++ b/src/conductor/config/schema.py @@ -763,7 +763,16 @@ class AgentDef(BaseModel): """Environment variables for script subprocess.""" working_dir: str | None = None - """Working directory for script subprocess execution.""" + """Working directory for the script subprocess OR a provider-backed agent + session and its MCP servers. + + On ``type: script`` steps it sets the subprocess cwd. On provider-backed + LLM agents it is resolved by the engine (Jinja-rendered, then relative + paths resolve against the workflow file's directory) and applied to the + provider session cwd and all of the agent's stdio MCP servers. Falls back + to ``runtime.working_dir`` when unset on the agent. Rejected on + wait/set/terminate/human_gate/workflow step types. + """ stdin: str | None = None """Payload written to the script subprocess's stdin (script type only). @@ -1180,6 +1189,8 @@ def validate_agent_type(self) -> AgentDef: ) if self.output_mode is not None: raise ValueError("human_gate agents cannot have 'output_mode'") + if self.working_dir: + raise ValueError("human_gate agents cannot have 'working_dir'") elif self.type == "script": if not self.command: raise ValueError("script agents require 'command'") @@ -1263,6 +1274,8 @@ def validate_agent_type(self) -> AgentDef: raise ValueError("workflow agents cannot have 'output_type' (only 'set' agents do)") if self.output_mode is not None: raise ValueError("workflow agents cannot have 'output_mode'") + if self.working_dir: + raise ValueError("workflow agents cannot have 'working_dir'") elif self.type == "wait": if self.duration is None: raise ValueError("wait agents require 'duration'") @@ -2072,6 +2085,16 @@ def _coerce_provider(cls, value: Any) -> Any: ``create_session`` ``context_tier`` param). Other providers ignore it. """ + working_dir: str | None = None + """Workflow-wide default working directory for provider-backed agents. + + Acts as the fallback for every LLM agent that does not set its own + ``working_dir`` (agent value wins). Supports Jinja2 templating and is + resolved by the engine against the workflow file's directory before + reaching the provider. ``conductor validate`` errors when the resolved + provider declares ``capabilities.working_dir=False``. + """ + class WorkflowDef(BaseModel): """Top-level workflow configuration.""" diff --git a/src/conductor/config/validator.py b/src/conductor/config/validator.py index 81b52bcb..9be0168c 100644 --- a/src/conductor/config/validator.py +++ b/src/conductor/config/validator.py @@ -1582,6 +1582,7 @@ def _validate_provider_capabilities( # no matter where future call sites are added. runtime_default_effort = config.workflow.runtime.default_reasoning_effort runtime_max_session_seconds = config.workflow.runtime.max_session_seconds + runtime_working_dir = config.workflow.runtime.working_dir # Cache per provider name so we don't re-resolve for every agent. cache: dict[str, ProviderCapabilities] = {} @@ -1730,6 +1731,17 @@ def _check_agent_capabilities( f"(capabilities.max_session_seconds=False)." ) + # working_dir: a provider that cannot apply the directory would + # silently run the agent (and its MCP servers) in the wrong cwd — + # the same class of silently-dropped operational intent as + # max_session_seconds. + if agent.working_dir is not None and not caps.working_dir: + errors.append( + f"Agent '{agent.name}' sets working_dir={agent.working_dir!r} " + f"but provider '{provider_name}' does not apply agent working " + f"directories (capabilities.working_dir=False)." + ) + # All provider-backed agents that run at workflow scope: top-level agents # PLUS for_each inline agents (``ForEachDef.agent``), which inherit the # workflow-level ``mcp_servers`` / ``max_session_seconds`` and run with @@ -1795,6 +1807,32 @@ def _check_agent_capabilities( f"max_session_seconds." ) + # ----- Workflow-level: working_dir ----- + # A runtime-wide working_dir is inherited by every LLM agent that does + # not set its own. A provider that cannot apply it would silently run + # those agents in the wrong directory — error against every resolved + # provider that actually receives the setting. + if runtime_working_dir is not None: + providers_inheriting_working_dir: dict[str, list[str]] = {} + for agent in all_llm_agents: + # A per-agent working_dir overrides the runtime default; that + # case is checked in ``_check_agent_capabilities`` instead. + if agent.working_dir is not None: + continue + pname = _resolved_provider_name(agent, default_provider) + providers_inheriting_working_dir.setdefault(pname, []).append(agent.name) + for pname, agent_names in providers_inheriting_working_dir.items(): + pcaps = _caps_for(pname) + if pcaps is not None and not pcaps.working_dir: + errors.append( + f"Workflow declares 'runtime.working_dir'={runtime_working_dir!r} " + f"but provider '{pname}' does not apply agent working directories " + f"(capabilities.working_dir=False) and is used by agent(s): " + f"{sorted(agent_names)!r}. Override these agents to a provider " + f"with working-directory support, or remove the workflow-level " + f"working_dir." + ) + # ----- Per-agent checks ----- for agent in config.agents: if not _is_llm_agent(agent): diff --git a/src/conductor/engine/checkpoint.py b/src/conductor/engine/checkpoint.py index ce733a1f..2502aff0 100644 --- a/src/conductor/engine/checkpoint.py +++ b/src/conductor/engine/checkpoint.py @@ -89,6 +89,11 @@ class CheckpointData: context: Serialized ``WorkflowContext`` state. limits: Serialized ``LimitEnforcer`` state. copilot_session_ids: Mapping of agent names to Copilot session IDs. + copilot_session_cwds: Mapping of agent names to the working directory + their Copilot session was created with. Persisted so resume can + detect a changed cwd and start a fresh session instead of + resuming into the wrong directory. Empty for checkpoints written + by a version of Conductor that predated this field. file_path: Path where the checkpoint file is stored. instructions_preamble: Workspace instructions preamble that was active during the original run, or ``None``. @@ -116,6 +121,7 @@ class CheckpointData: context: dict[str, Any] limits: dict[str, Any] copilot_session_ids: dict[str, str] = field(default_factory=dict) + copilot_session_cwds: dict[str, str] = field(default_factory=dict) file_path: Path = field(default_factory=lambda: Path()) instructions_preamble: str | None = None """Workspace instructions preamble that was active during the original run.""" @@ -176,6 +182,7 @@ def save_checkpoint( error: BaseException | None, inputs: dict[str, Any], copilot_session_ids: dict[str, str] | None = None, + copilot_session_cwds: dict[str, str] | None = None, system_metadata: dict[str, Any] | None = None, instructions_preamble: str | None = None, run_id: str = "", @@ -201,6 +208,10 @@ def save_checkpoint( for a periodic / non-failure checkpoint. inputs: Workflow inputs. copilot_session_ids: Optional mapping of agent names to session IDs. + copilot_session_cwds: Optional mapping of agent names to the + working directory their session was created with. Persisted + alongside the session IDs so resume can skip stale sessions + whose cwd no longer matches the agent's resolved cwd. system_metadata: Optional system metadata captured at workflow start. instructions_preamble: Optional workspace instructions preamble to persist. run_id: Original run identifier (from ``EventLogSubscriber``). @@ -255,6 +266,7 @@ def save_checkpoint( "context": _make_json_serializable(context.to_dict()), "limits": _make_json_serializable(limits.to_dict()), "copilot_session_ids": copilot_session_ids or {}, + "copilot_session_cwds": copilot_session_cwds or {}, "system": system_metadata or {}, "instructions_preamble": instructions_preamble, "run_id": run_id, @@ -373,6 +385,9 @@ def load_checkpoint(checkpoint_path: Path) -> CheckpointData: context=data["context"], limits=data["limits"], copilot_session_ids=data.get("copilot_session_ids", {}), + # Backward compatible: pre-cwd checkpoints have no such key and + # load with an empty mapping (legacy resume-by-id behavior). + copilot_session_cwds=data.get("copilot_session_cwds", {}), file_path=checkpoint_path, instructions_preamble=data.get("instructions_preamble"), run_id=data.get("run_id", "") or "", diff --git a/src/conductor/engine/validator.py b/src/conductor/engine/validator.py index dafe648b..44b63044 100644 --- a/src/conductor/engine/validator.py +++ b/src/conductor/engine/validator.py @@ -218,6 +218,7 @@ def _build_validator_agent(self, agent: AgentDef) -> AgentDef: system_prompt=VALIDATOR_SYSTEM_PROMPT.format(criteria=agent.validator.criteria), tools=[], output=_VALIDATOR_OUTPUT_SCHEMA, + working_dir=agent.working_dir, ) def _parse(self, content: Any) -> tuple[bool, list[str], bool]: diff --git a/src/conductor/engine/workflow.py b/src/conductor/engine/workflow.py index e035c876..ab3a316a 100644 --- a/src/conductor/engine/workflow.py +++ b/src/conductor/engine/workflow.py @@ -11,6 +11,7 @@ import copy import json import logging +import os import sys import time as _time import uuid @@ -484,6 +485,48 @@ def _workflow_dir(self) -> Path | None: """Resolved parent directory of the workflow file, or None if unset.""" return Path(self.workflow_path).resolve().parent if self.workflow_path else None + def _resolve_agent_working_dir( + self, agent: AgentDef, agent_context: dict[str, Any] + ) -> AgentDef: + """Resolve an agent's effective ``working_dir`` and return an updated copy. + + Precedence is ``agent.working_dir`` over ``runtime.working_dir``; the + chosen raw value is Jinja-rendered against the per-agent context (so + both levels support templates, e.g. ``{{ item }}`` in for-each), then + ``~``-expanded, made absolute against the workflow file's directory + (falling back to the process cwd), and lexically normalised with + :func:`os.path.normpath` (``resolve()`` is deliberately not used so + symlink aliases stay distinct). A missing directory raises + :class:`ExecutionError` before any provider call. When neither level + sets a value the agent is returned unchanged (``working_dir=None`` and + the provider uses its own cwd). + """ + raw = agent.working_dir + if raw is None: + raw = self.config.workflow.runtime.working_dir + if raw is None: + return agent + + rendered = self.renderer.render(raw, agent_context) + path = Path(rendered).expanduser() + if not path.is_absolute(): + base = self._workflow_dir if self._workflow_dir is not None else Path.cwd() + path = base / path + resolved = os.path.normpath(path) + + if not Path(resolved).is_dir(): + raise ExecutionError( + f"Agent '{agent.name}': working_dir '{resolved}' does not exist or " + f"is not a directory (rendered from '{raw}')", + agent_name=agent.name, + suggestion=( + "Create the directory before the agent runs (e.g. via a " + "script step) or fix the working_dir template." + ), + ) + + return agent.model_copy(update={"working_dir": resolved}) + def _build_pricing_overrides(self) -> dict[str, ModelPricing] | None: """Build pricing overrides from workflow cost configuration. @@ -1911,18 +1954,24 @@ def _write_checkpoint( # provider raising here must not break the (failure or periodic) # checkpoint save, so the "never raises" contract holds for both paths. copilot_session_ids: dict[str, str] | None = None + copilot_session_cwds: dict[str, str] | None = None try: provider = self._single_provider if provider is not None and hasattr(provider, "get_session_ids"): copilot_session_ids = provider.get_session_ids() # type: ignore[union-attr] + if hasattr(provider, "get_session_cwds"): + copilot_session_cwds = provider.get_session_cwds() # type: ignore[union-attr] elif self._registry is not None: for p in self._registry.get_active_providers().values(): if hasattr(p, "get_session_ids"): copilot_session_ids = p.get_session_ids() # type: ignore[union-attr] + if hasattr(p, "get_session_cwds"): + copilot_session_cwds = p.get_session_cwds() # type: ignore[union-attr] break except Exception: logger.warning("Failed to collect provider session IDs for checkpoint", exc_info=True) copilot_session_ids = None + copilot_session_cwds = None return CheckpointManager.save_checkpoint( workflow_path=self.workflow_path, @@ -1932,6 +1981,7 @@ def _write_checkpoint( error=error, inputs=self.context.workflow_inputs, copilot_session_ids=copilot_session_ids, + copilot_session_cwds=copilot_session_cwds, system_metadata=self._system_metadata, instructions_preamble=self._instructions_preamble, run_id=self._run_context.run_id, @@ -3073,20 +3123,47 @@ async def _execute_loop(self, current_agent_name: str) -> dict[str, Any]: self.limits.get_agent_execution_count(agent.name) + 1 ) - self._emit( - "agent_started", - { - "agent_name": agent.name, - "iteration": agent_execution_count, - "agent_type": agent.type or "agent", - "context_window_max": await self._get_context_window_for_agent( - agent - ), - }, + # Trim context BEFORE building the per-agent context and + # BEFORE agent_started so a max_tokens workflow emits the + # start event with the prompt already trimmed (issue: + # working_dir ordering). Trim is context-local — nothing + # between the old (post-agent_started) and new position + # reads the untrimmed context. + self._trim_context_if_needed() + + # Build the per-agent context once for every step type; + # the type-specific branches below reuse it instead of + # re-building (human_gate still prefers the full-template + # context and overwrites this local). + agent_context = self.context.build_for_agent( + agent.name, + agent.input, + mode=self.config.workflow.context.mode, + agent_type=agent.type, ) - # Trim context if max_tokens is configured - self._trim_context_if_needed() + # Resolve working_dir only for provider-backed LLM agents + # (type None/"agent"). wait/set/terminate/human_gate/ + # workflow are schema-rejected from declaring one, and + # script resolves its own in ScriptExecutor. + is_llm_agent = agent.type in (None, "agent") + resolved_agent = ( + self._resolve_agent_working_dir(agent, agent_context) + if is_llm_agent + else agent + ) + + started_payload: dict[str, Any] = { + "agent_name": agent.name, + "iteration": agent_execution_count, + "agent_type": agent.type or "agent", + "context_window_max": await self._get_context_window_for_agent( + resolved_agent + ), + } + if is_llm_agent: + started_payload["working_dir"] = resolved_agent.working_dir + self._emit("agent_started", started_payload) # Handle terminate steps — explicit workflow exit with a # structured reason and status. Reached via a normal @@ -3097,12 +3174,6 @@ async def _execute_loop(self, current_agent_name: str) -> dict[str, Any]: # evaluated after. if agent.type == "terminate": terminate_elapsed = _time.time() - _workflow_start - agent_context = self.context.build_for_agent( - agent.name, - agent.input, - mode=self.config.workflow.context.mode, - agent_type=agent.type, - ) # Render the reason against context first so the # rendered value is available to output_template # and to the workflow-level output: fallback. @@ -3289,12 +3360,6 @@ async def _execute_loop(self, current_agent_name: str) -> dict[str, Any]: # Handle script steps if agent.type == "script": - agent_context = self.context.build_for_agent( - agent.name, - agent.input, - mode=self.config.workflow.context.mode, - agent_type=agent.type, - ) _script_start = _time.time() # Count how many times this specific script has been executed @@ -3438,12 +3503,6 @@ async def _execute_loop(self, current_agent_name: str) -> dict[str, Any]: # Handle wait steps if agent.type == "wait": - agent_context = self.context.build_for_agent( - agent.name, - agent.input, - mode=self.config.workflow.context.mode, - agent_type=agent.type, - ) _wait_start = _time.time() wait_execution_count = ( @@ -3580,13 +3639,6 @@ async def _execute_loop(self, current_agent_name: str) -> dict[str, Any]: # Handle set steps. Pure context transformations: # render, coerce, validate, emit, route. if agent.type == "set": - agent_context = self.context.build_for_agent( - agent.name, - agent.input, - mode=self.config.workflow.context.mode, - agent_type=agent.type, - ) - set_output = await self._run_set_step(agent, agent_context) self.context.store(agent.name, set_output.value) self.limits.record_execution(agent.name) @@ -3635,12 +3687,6 @@ async def _execute_loop(self, current_agent_name: str) -> dict[str, Any]: # Handle sub-workflow steps if agent.type == "workflow": - agent_context = self.context.build_for_agent( - agent.name, - agent.input, - mode=self.config.workflow.context.mode, - agent_type=agent.type, - ) _sub_start = _time.time() sub_execution_count = ( @@ -3725,23 +3771,20 @@ async def _execute_loop(self, current_agent_name: str) -> dict[str, Any]: ) continue - # Build context for this agent - agent_context = self.context.build_for_agent( - agent.name, - agent.input, - mode=self.config.workflow.context.mode, - agent_type=agent.type, - ) - - # Execute agent (get executor for multi-provider support) + # Execute agent (get executor for multi-provider support). + # resolved_agent carries the engine-resolved working_dir; + # agent_context was built (and trimmed) above, before + # agent_started. Subsequent model_copy(update={...}) calls + # inside AgentExecutor merge, so the resolved working_dir + # survives to the provider. _agent_start = _time.time() - executor = await self._get_executor_for_agent(agent) + executor = await self._get_executor_for_agent(resolved_agent) guidance_section = self.context.get_guidance_prompt_section() event_callback = self._make_event_callback(agent.name) output = await self._execute_with_agent_timeout( - agent, + resolved_agent, executor.execute( - agent, + resolved_agent, agent_context, guidance_section=guidance_section, interrupt_signal=self._interrupt_event, @@ -3771,7 +3814,7 @@ async def _execute_loop(self, current_agent_name: str) -> dict[str, Any]: self._interrupt_event.clear() continue output = await self._handle_partial_output( - agent, + resolved_agent, output, agent_context, guidance_section, @@ -3783,7 +3826,7 @@ async def _execute_loop(self, current_agent_name: str) -> dict[str, Any]: # Dialog mode: evaluate whether agent should enter dialog if agent.dialog and not output.partial: output = await self._handle_dialog( - agent, + resolved_agent, output, agent_context, executor, @@ -3793,7 +3836,7 @@ async def _execute_loop(self, current_agent_name: str) -> dict[str, Any]: # Validator: grade output and re-run once on failure if agent.validator and not output.partial: output = await self._apply_validator( - agent, + resolved_agent, output, _agent_elapsed, agent_context, @@ -3804,7 +3847,7 @@ async def _execute_loop(self, current_agent_name: str) -> dict[str, Any]: _agent_elapsed = _time.time() - _agent_start # Record usage and calculate cost - await self._ensure_pricing_resolved(agent, output.model) + await self._ensure_pricing_resolved(resolved_agent, output.model) usage = self.usage_tracker.record(agent.name, output, _agent_elapsed) output_keys = ( @@ -3825,7 +3868,7 @@ async def _execute_loop(self, current_agent_name: str) -> dict[str, Any]: "output_keys": output_keys, "context_window_used": output.input_tokens, "context_window_max": await self._get_context_window_for_agent( - agent, output + resolved_agent, output ), }, ) @@ -3841,7 +3884,7 @@ async def _execute_loop(self, current_agent_name: str) -> dict[str, Any]: self._check_budget() # Evaluate routes using the Router - route_result = self._evaluate_routes(agent, output.content) + route_result = self._evaluate_routes(resolved_agent, output.content) self._emit( "route_taken", @@ -4836,13 +4879,30 @@ async def execute_single_agent(agent: AgentDef) -> tuple[str, Any]: ) return (agent.name, set_output.value) + # Resolve working_dir for provider-backed LLM agents against + # this agent's own (pre-group snapshot) context. `set` steps + # returned above; other types in a parallel group are LLM agents. + resolved_agent = self._resolve_agent_working_dir(agent, agent_context) + + # LLM-only per-member start event: emitted only here (after the + # per-agent resolution) so ``working_dir`` is the resolved value; + # the pre-context envelope ``parallel_started`` stays unchanged. + self._emit( + "parallel_agent_started", + { + "group_name": parallel_group.name, + "agent_name": agent.name, + "working_dir": resolved_agent.working_dir, + }, + ) + # Execute agent (get executor for multi-provider support) - executor = await self._get_executor_for_agent(agent) + executor = await self._get_executor_for_agent(resolved_agent) event_callback = self._make_event_callback(agent.name) output = await self._execute_with_agent_timeout( - agent, + resolved_agent, executor.execute( - agent, + resolved_agent, agent_context, event_callback=event_callback, ), @@ -4850,9 +4910,9 @@ async def execute_single_agent(agent: AgentDef) -> tuple[str, Any]: _agent_elapsed = _time.time() - _agent_start # Validator: grade output and re-run once on failure - if agent.validator and not output.partial: + if resolved_agent.validator and not output.partial: output = await self._apply_validator( - agent, + resolved_agent, output, _agent_elapsed, agent_context, @@ -4863,7 +4923,7 @@ async def execute_single_agent(agent: AgentDef) -> tuple[str, Any]: _agent_elapsed = _time.time() - _agent_start # Record usage and calculate cost - await self._ensure_pricing_resolved(agent, output.model) + await self._ensure_pricing_resolved(resolved_agent, output.model) usage = self.usage_tracker.record(agent.name, output, _agent_elapsed) self._emit( @@ -4877,7 +4937,7 @@ async def execute_single_agent(agent: AgentDef) -> tuple[str, Any]: "cost_usd": usage.cost_usd, "context_window_used": output.input_tokens, "context_window_max": await self._get_context_window_for_agent( - agent, output + resolved_agent, output ), }, ) @@ -5293,8 +5353,6 @@ async def execute_single_item(item: Any, index: int, key: str) -> tuple[str, Any ) return (key, set_output.value) - executor = await self._get_executor_for_agent(for_each_group.agent) - # Qualify the per-iteration agent name so that any verbose # provider-side logging (e.g. CopilotProvider tool/reasoning # lines) can attribute interleaved output to a specific @@ -5304,6 +5362,27 @@ async def execute_single_item(item: Any, index: int, key: str) -> tuple[str, Any update={"name": f"{for_each_group.agent.name}[{key}]"} ) + # Resolve working_dir AFTER loop variables were injected into + # agent_context so a `{{ item }}` (or `{{ }}`) template in + # the path resolves to this iteration's value. + qualified_agent = self._resolve_agent_working_dir(qualified_agent, agent_context) + + # LLM-only per-item start event: emitted only here (after the + # per-item resolution) so ``working_dir`` is the resolved value; + # the pre-context envelope ``for_each_item_started`` stays + # unchanged. + self._emit( + "for_each_agent_started", + { + "group_name": for_each_group.name, + "agent_name": qualified_agent.name, + "item_key": key, + "working_dir": qualified_agent.working_dir, + }, + ) + + executor = await self._get_executor_for_agent(qualified_agent) + # Item-scoped event callback that tags all streaming events # with the for-each group name + item_key. Wrapper keys are # placed *after* ``**data`` so they override any qualified diff --git a/src/conductor/mcp/manager.py b/src/conductor/mcp/manager.py index d7914050..50bdeee2 100644 --- a/src/conductor/mcp/manager.py +++ b/src/conductor/mcp/manager.py @@ -77,6 +77,7 @@ async def connect_server( args: list[str] | None = None, env: dict[str, str] | None = None, timeout: int | None = None, + cwd: str | None = None, ) -> list[dict[str, Any]]: """Connect to an MCP server and return its tools. @@ -89,6 +90,9 @@ async def connect_server( args: Command arguments. env: Environment variables for the server process. timeout: Connection timeout in seconds (not currently used). + cwd: Working directory for the spawned server process. When None, + the server inherits the conductor process's current working + directory (pre-pool legacy behavior). Returns: List of tool definitions from this server. Each tool dict contains: @@ -111,6 +115,7 @@ async def connect_server( command=command, args=args or [], env=env, + cwd=cwd, ) try: diff --git a/src/conductor/providers/capabilities.py b/src/conductor/providers/capabilities.py index 1ab267d8..2febbb49 100644 --- a/src/conductor/providers/capabilities.py +++ b/src/conductor/providers/capabilities.py @@ -130,6 +130,14 @@ class ProviderCapabilities(BaseModel): :class:`ParallelGroup`, and may only appear in a :class:`ForEachDef` whose ``max_concurrent`` is 1.""" + working_dir: bool = False + """``True`` when the provider applies an agent's resolved ``working_dir`` + to the SDK session cwd and the agent's stdio MCP servers. Workflows that + set ``working_dir`` (per-agent or via ``runtime.working_dir``) against a + provider with ``working_dir=False`` fail validation — silently ignoring + the directory would run the agent in the wrong repository. Defaults to + ``False`` (conservative).""" + upstream_pin: str | None = None """Upstream package pin surfaced in the experimental banner, e.g. ``"claude-agent-sdk>=0.1.0"``. ``None`` for providers that have no @@ -198,6 +206,8 @@ def declared_limitations(self) -> list[str]: items.append("no usage tracking") if not self.concurrent_safe: items.append("not safe to run in parallel") + if not self.working_dir: + items.append("working_dir ignored") return items @@ -249,6 +259,7 @@ def _build_unimplemented_placeholder() -> ProviderCapabilities: checkpoint_resume=True, usage_tracking=True, concurrent_safe=True, + working_dir=True, upstream_pin=None, maintainer="(not yet implemented)", ) diff --git a/src/conductor/providers/claude.py b/src/conductor/providers/claude.py index 8589e3fb..c7083f10 100644 --- a/src/conductor/providers/claude.py +++ b/src/conductor/providers/claude.py @@ -22,6 +22,7 @@ import contextlib import json import logging +import os import random import time from typing import TYPE_CHECKING, Any, Protocol, get_args @@ -164,6 +165,10 @@ class ClaudeProvider(AgentProvider): usage_tracking=True, # No global mutable state — safe to run N parallel agents. concurrent_safe=True, + # The resolved ``working_dir`` selects the MCPManager pool key and is + # forwarded to ``StdioServerParameters(cwd=...)`` for each stdio MCP + # server the agent connects. + working_dir=True, upstream_pin=None, maintainer="@microsoft/conductor", ) @@ -252,9 +257,17 @@ def __init__( self._default_max_session_seconds = max_session_seconds self._default_reasoning_effort: ReasoningEffort | None = default_reasoning_effort - # MCP server configuration for tool support + # MCP server configuration for tool support. + # Managers are pooled by resolved working directory: each distinct + # cwd gets its own MCPManager because stdio MCP servers are spawned + # with the manager's cwd. The pool lifecycle is bounded by the + # provider lifetime — close() shuts down every pooled manager. v1 + # intentionally has no eviction/LRU: the number of distinct cwds in + # a workflow run is expected to be small, and evicting a live + # manager would kill in-flight tool calls. self._mcp_servers_config = mcp_servers - self._mcp_manager: MCPManager | None = None + self._mcp_managers: dict[str, MCPManager] = {} + self._mcp_manager_locks: dict[str, asyncio.Lock] = {} # Cache of model_id -> max_input_tokens populated lazily on first # get_max_prompt_tokens() call. Guarded by an asyncio.Lock to avoid @@ -560,18 +573,40 @@ async def get_model_capabilities(self, model: str) -> ModelCapabilityInfo | None max_context_window_tokens=None, ) - async def _ensure_mcp_connected(self) -> None: - """Connect to MCP servers if configured. + async def _get_mcp_manager_for_cwd(self, resolved_cwd: str) -> MCPManager | None: + """Return the pooled MCPManager for ``resolved_cwd``, connecting on first use. - This method lazily initializes MCP connections on first use. - It creates an MCPManager instance and connects to all configured - MCP servers. + Each distinct working directory gets its own MCPManager so stdio MCP + servers are spawned with that directory as their ``cwd``. The lazy + connect is guarded by a per-cwd ``asyncio.Lock`` so parallel agents + resolving the same cwd observe exactly one manager (no duplicate + spawns), while agents with different cwds proceed concurrently. + + Per-server connect is fail-open: a server that fails to connect is + logged and skipped, and the manager is still pooled as long as at + least one server connected. When NO servers connect the manager is + returned but not pooled, so the next agent for the same cwd retries + the connect (a transient spawn failure does not become permanent). + + Pool lifecycle is bounded by the provider lifetime — ``close()`` + shuts down every pooled manager. v1 intentionally has no + eviction/LRU. + + Args: + resolved_cwd: Absolute, normalized working directory that keys + the pool (``agent.working_dir or os.getcwd()`` at the call + site; the engine has already resolved ``agent.working_dir`` + to an absolute normpath). + + Returns: + The pooled manager, or None when no MCP servers are configured + or the MCP SDK is not installed. """ - # Skip if already initialized or no servers configured - if self._mcp_manager is not None: - return + # Fast path: already pooled. + if resolved_cwd in self._mcp_managers: + return self._mcp_managers[resolved_cwd] if not self._mcp_servers_config: - return + return None from conductor.mcp.manager import MCP_SDK_AVAILABLE, MCPManager @@ -580,49 +615,78 @@ async def _ensure_mcp_connected(self) -> None: "MCP servers configured but MCP SDK not installed. " "Install with: uv add 'mcp>=1.0.0'" ) - return - - self._mcp_manager = MCPManager() + return None - for name, config in self._mcp_servers_config.items(): - server_type = config.get("type", "stdio") - if server_type == "stdio": - try: - await self._mcp_manager.connect_server( - name=name, - command=config["command"], - args=config.get("args", []), - env=config.get("env"), - timeout=config.get("timeout"), + # No guard needed around lock creation: there is no await between the + # fast-path check above and the per-cwd lock acquisition below, so + # concurrent coroutines cannot interleave here. + lock = self._mcp_manager_locks.get(resolved_cwd) + if lock is None: + lock = asyncio.Lock() + self._mcp_manager_locks[resolved_cwd] = lock + + async with lock: + # Re-check under the per-cwd lock: a concurrent agent may have + # connected while we were waiting. + if resolved_cwd in self._mcp_managers: + return self._mcp_managers[resolved_cwd] + + manager = MCPManager() + for name, config in self._mcp_servers_config.items(): + server_type = config.get("type", "stdio") + if server_type == "stdio": + try: + await manager.connect_server( + name=name, + command=config["command"], + args=config.get("args", []), + env=config.get("env"), + timeout=config.get("timeout"), + cwd=resolved_cwd, + ) + logger.info(f"Connected to MCP server '{name}' (cwd={resolved_cwd})") + except Exception as e: + logger.error(f"Failed to connect to MCP server '{name}': {e}") + # Continue with other servers (fail-open per server) + else: + logger.warning( + f"MCP server '{name}' has unsupported type '{server_type}' " + "(Claude provider only supports 'stdio')" ) - logger.info(f"Connected to MCP server '{name}'") - except Exception as e: - logger.error(f"Failed to connect to MCP server '{name}': {e}") - # Continue with other servers + + if manager.has_servers(): + self._mcp_managers[resolved_cwd] = manager else: logger.warning( - f"MCP server '{name}' has unsupported type '{server_type}' " - "(Claude provider only supports 'stdio')" + "No MCP servers connected for cwd=%s; manager not pooled so " + "the next agent for this cwd will retry the connect.", + resolved_cwd, ) + return manager def _convert_mcp_tools_to_claude( self, tool_filter: list[str] | None = None, + manager: MCPManager | None = None, ) -> list[dict[str, Any]]: """Convert MCP tools to Claude tool format. Args: tool_filter: Optional list of tool names to include (prefixed names). If None, all tools are included. + manager: The pooled MCPManager for this agent's working directory. + The manager is passed explicitly (never read from shared + mutable provider state) so parallel agents with different + cwds cannot observe each other's tools. Returns: List of tool definitions in Claude's expected format. """ - if not self._mcp_manager: + if not manager: return [] claude_tools: list[dict[str, Any]] = [] - for tool in self._mcp_manager.get_all_tools(): + for tool in manager.get_all_tools(): # Apply filter if specified if tool_filter and tool["name"] not in tool_filter: continue @@ -638,12 +702,21 @@ def _convert_mcp_tools_to_claude( return claude_tools async def close(self) -> None: - """Release provider resources and close connections.""" - # Close MCP connections first - if self._mcp_manager is not None: - await self._mcp_manager.close() - self._mcp_manager = None - logger.debug("MCP manager closed") + """Release provider resources and close connections. + + Shuts down every pooled MCPManager (one per distinct working + directory). Idempotent: a second call is a no-op. + """ + # Close MCP connections first (all pool entries). + if self._mcp_managers: + for cwd, manager in self._mcp_managers.items(): + try: + await manager.close() + except Exception as e: + logger.warning(f"Error closing MCP manager for cwd={cwd}: {e}") + self._mcp_managers.clear() + self._mcp_manager_locks.clear() + logger.debug("All pooled MCP managers closed") if self._client is not None: # Drop the client reference *before* awaiting close() so any @@ -1020,8 +1093,12 @@ async def _execute_with_retry( if self._client is None: raise ProviderError("Claude client not initialized") - # Connect to MCP servers if configured (lazy initialization) - await self._ensure_mcp_connected() + # Resolve this agent's MCP manager from the cwd pool (lazy connect). + # The manager is a LOCAL variable threaded through the whole agentic + # loop — never stored as shared mutable provider state — so parallel + # agents with different working directories stay isolated. + resolved_cwd = agent.working_dir or os.getcwd() + mcp_manager = await self._get_mcp_manager_for_cwd(resolved_cwd) last_error: Exception | None = None config = self._resolve_retry_config(agent) @@ -1093,8 +1170,8 @@ async def _execute_with_retry( ) # Add MCP tools if available - if self._mcp_manager and self._mcp_manager.has_servers(): - mcp_tools = self._convert_mcp_tools_to_claude(tools) # tools is the filter + if mcp_manager and mcp_manager.has_servers(): + mcp_tools = self._convert_mcp_tools_to_claude(tools, mcp_manager) # tools is the filter all_tools.extend(mcp_tools) if mcp_tools: logger.debug(f"Added {len(mcp_tools)} MCP tools to request") @@ -1120,6 +1197,7 @@ async def _execute_with_retry( thinking=thinking, max_parse_recovery_attempts=config.max_parse_recovery_attempts, system_prompt=system_prompt, + mcp_manager=mcp_manager, ) # Handle partial output from mid-agent interrupt @@ -1476,6 +1554,7 @@ async def _execute_agentic_loop( thinking: dict[str, Any] | None = None, max_parse_recovery_attempts: int | None = None, system_prompt: str | None = None, + mcp_manager: MCPManager | None = None, ) -> tuple[ClaudeResponse, int | None, bool]: """Execute an agentic loop that handles MCP tool calls. @@ -1506,6 +1585,10 @@ async def _execute_agentic_loop( max_parse_recovery_attempts: Resolved per-agent parse recovery limit. None means use the provider-level default. system_prompt: Optional rendered system prompt forwarded to every API call. + mcp_manager: Pooled MCPManager for this agent's working directory. + Passed explicitly (never read from shared provider state) so + parallel agents with different cwds execute their tool calls + against the correct per-cwd connection pool. Returns: Tuple of (final_response, total_tokens_used, is_partial). @@ -1708,7 +1791,7 @@ async def _execute_agentic_loop( # No MCP tools to execute return response, total_tokens, False - if not self._mcp_manager: + if not mcp_manager: logger.warning( f"Claude called MCP tools but no MCP manager available: " f"{[t.name for t in mcp_tool_uses]}" @@ -1739,7 +1822,7 @@ async def _execute_agentic_loop( logger.debug("Error in event_callback for agent_tool_start", exc_info=True) try: - result = await self._mcp_manager.call_tool( + result = await mcp_manager.call_tool( tool_use.name, dict(tool_use.input) if hasattr(tool_use, "input") else {} ) tool_results.append( diff --git a/src/conductor/providers/claude_agent_sdk.py b/src/conductor/providers/claude_agent_sdk.py index 61f7a41c..566df5ae 100644 --- a/src/conductor/providers/claude_agent_sdk.py +++ b/src/conductor/providers/claude_agent_sdk.py @@ -159,6 +159,11 @@ class ClaudeAgentSdkProvider(AgentProvider): # No global mutable state shared across calls — the SDK spawns # an independent subprocess per query() invocation. concurrent_safe=True, + # MCP servers are rejected at the factory (mcp_tools=False), so + # there is nothing to scope by directory — the SDK CLI manages its + # own cwd. Declared False so ``conductor validate`` errors when a + # workflow sets ``working_dir`` against this provider. + working_dir=False, upstream_pin="claude-agent-sdk>=0.1.0", maintainer="@lesandiz (best-effort)", ) diff --git a/src/conductor/providers/copilot.py b/src/conductor/providers/copilot.py index 1a1a2650..d9039380 100644 --- a/src/conductor/providers/copilot.py +++ b/src/conductor/providers/copilot.py @@ -231,6 +231,10 @@ class CopilotProvider(AgentProvider): usage_tracking=True, # No global mutable state — safe to run N parallel agents. concurrent_safe=True, + # The resolved ``working_dir`` is applied to the SDK session's + # ``working_directory`` and stamped onto each stdio MCP server + # config per execution (no shared-dict mutation). + working_dir=True, upstream_pin=None, maintainer="@microsoft/conductor", ) @@ -305,6 +309,8 @@ def __init__( self._max_schema_depth = 10 # Max nesting depth for recursive schema building self._session_ids: dict[str, str] = {} self._resume_session_ids: dict[str, str] = {} + self._session_cwds: dict[str, str] = {} + self._resume_session_cwds: dict[str, str] = {} self._interrupted_session: Any = None self._abort_supported: bool | None = None self._provider_settings = provider_settings @@ -790,6 +796,12 @@ async def _execute_sdk_call( ) try: + # Resolve the session working directory: the engine resolves + # ``agent.working_dir`` (Jinja render, absolutize, is_dir check) + # before dispatching to the provider; ``None`` keeps the legacy + # process-cwd behavior. + resolved_cwd = agent.working_dir or os.getcwd() + # Build session kwargs for the SDK. # # ``streaming=True`` is required: in non-streaming mode the model @@ -806,7 +818,7 @@ async def _execute_sdk_call( session_kwargs: dict[str, Any] = { "model": model, "on_permission_request": self._default_permission_handler, - "working_directory": os.getcwd(), + "working_directory": resolved_cwd, "streaming": True, } @@ -820,9 +832,12 @@ async def _execute_sdk_call( self._temperature, ) - # Add MCP servers if configured + # Add MCP servers if configured. Stdio/local servers get a + # per-execution copy stamped with the resolved working directory; + # the shared ``self._mcp_servers`` mapping is never mutated and + # remote (http/sse) servers are left untouched. if self._mcp_servers: - session_kwargs["mcp_servers"] = self._mcp_servers + session_kwargs["mcp_servers"] = self._mcp_servers_for_cwd(resolved_cwd) # Apply custom provider routing (Ollama / vLLM / Azure / etc.) # when runtime.provider opted into it. @@ -857,31 +872,67 @@ async def _execute_sdk_call( model, ) - # Attempt to resume a previous session if one exists for this agent + # Attempt to resume a previous session if one exists for this agent. + # Resume is only valid when the session was originally created with + # the same working directory: the SDK bakes cwd into the session's + # workspace, so resuming under a different cwd would silently run + # the agent in the wrong directory. When the recorded cwd differs + # (or is unknown for pre-cwd checkpoints, where we keep the legacy + # resume-by-id behavior), fall through to ``create_session``. session: Any = None resume_sid = self._resume_session_ids.get(agent.name) if resume_sid is not None: - try: - session = await self._client.resume_session( + recorded_cwd = self._resume_session_cwds.get(agent.name) + should_resume = True + if recorded_cwd is None: + logger.info( + "Resuming Copilot session %s for agent '%s' without a " + "recorded working directory; checkpoint predates " + "working_dir tracking, preserving legacy resume-by-id " + "behavior with resolved cwd %s.", resume_sid, - on_permission_request=self._default_permission_handler, + agent.name, + resolved_cwd, ) - logger.info(f"Resumed Copilot session {resume_sid} for agent '{agent.name}'") - except Exception as exc: + elif recorded_cwd != resolved_cwd: logger.warning( - f"Could not resume session {resume_sid} for agent " - f"'{agent.name}': {exc}. Falling back to new session." + "Skipping resume of Copilot session %s for agent '%s': " + "working directory changed from %s to %s. Creating a new session.", + resume_sid, + agent.name, + recorded_cwd, + resolved_cwd, ) - session = None + should_resume = False + + if should_resume: + try: + resume_kwargs: dict[str, Any] = { + "on_permission_request": self._default_permission_handler, + "working_directory": resolved_cwd, + } + if self._mcp_servers: + resume_kwargs["mcp_servers"] = self._mcp_servers_for_cwd(resolved_cwd) + session = await self._client.resume_session(resume_sid, **resume_kwargs) + logger.info( + f"Resumed Copilot session {resume_sid} for agent '{agent.name}'" + ) + except Exception as exc: + logger.warning( + f"Could not resume session {resume_sid} for agent " + f"'{agent.name}': {exc}. Falling back to new session." + ) + session = None # Fall back to creating a new session if session is None: session = await self._client.create_session(**session_kwargs) - # Track session ID for checkpoint persistence + # Track session ID and resolved cwd for checkpoint persistence sid = getattr(session, "session_id", None) if sid is not None: self._session_ids[agent.name] = sid + self._session_cwds[agent.name] = resolved_cwd # Capture verbose state before callback (contextvars don't propagate to sync callbacks) from conductor.cli.app import is_full, is_verbose @@ -2567,6 +2618,33 @@ async def get_model_capabilities(self, model: str) -> ModelCapabilityInfo | None max_context_window_tokens=getattr(limits, "max_context_window_tokens", None), ) + def _mcp_servers_for_cwd(self, resolved_cwd: str) -> dict[str, Any]: + """Build a per-execution MCP server mapping stamped with ``resolved_cwd``. + + Stdio/local servers get a shallow-copied config with + ``working_directory`` set (the SDK translates it to the spawned + server's cwd). Remote (``http``/``sse``) servers are returned as-is — + a working directory is meaningless for a remote process. The shared + ``self._mcp_servers`` mapping and its nested dicts are never mutated, + so parallel agents with different cwds cannot race each other. + + Args: + resolved_cwd: Absolute working directory resolved by the engine + (or ``os.getcwd()`` when the agent declares no working_dir). + + Returns: + A new mapping of server name to config dict. + """ + stamped: dict[str, Any] = {} + for name, config in self._mcp_servers.items(): + if isinstance(config, dict) and config.get("type") in ("stdio", "local", None): + server_copy = dict(config) + server_copy["working_directory"] = resolved_cwd + stamped[name] = server_copy + else: + stamped[name] = config + return stamped + def get_session_ids(self) -> dict[str, str]: """Get tracked session IDs for all executed agents. @@ -2591,6 +2669,30 @@ def set_resume_session_ids(self, ids: dict[str, str]) -> None: """ self._resume_session_ids = dict(ids) + def get_session_cwds(self) -> dict[str, str]: + """Get the resolved working directory per executed agent. + + Mirrors :meth:`get_session_ids` — the engine persists this mapping in + the checkpoint (``copilot_session_cwds``) so a later resume can detect + when the resolved cwd changed since session creation and start a fresh + session instead of resuming into the wrong directory. + + Returns: + Dict mapping agent names to their session's working directory. + """ + return self._session_cwds.copy() + + def set_resume_session_cwds(self, cwds: dict[str, str]) -> None: + """Set session working directories restored from a checkpoint. + + Args: + cwds: Mapping of agent names to the working directory their + stored session was created with. Agents missing from this + mapping (pre-cwd checkpoints) keep the legacy resume-by-id + behavior. + """ + self._resume_session_cwds = dict(cwds) + def get_interrupted_session(self) -> Any | None: """Get the session handle kept alive after a mid-agent interrupt. diff --git a/src/conductor/providers/hermes.py b/src/conductor/providers/hermes.py index 2d7ff965..9a695486 100644 --- a/src/conductor/providers/hermes.py +++ b/src/conductor/providers/hermes.py @@ -92,6 +92,9 @@ class HermesProvider(AgentProvider): checkpoint_resume=True, usage_tracking=True, concurrent_safe=True, + # Hermes runs its own internal toolsets (mcp_tools=False); a + # per-agent working directory has no meaning for the session. + working_dir=False, upstream_pin="hermes-agent", maintainer="(community contribution)", ) diff --git a/src/conductor/providers/registry.py b/src/conductor/providers/registry.py index ad622a91..3edb3b05 100644 --- a/src/conductor/providers/registry.py +++ b/src/conductor/providers/registry.py @@ -53,6 +53,7 @@ def __init__( self._providers: dict[ProviderType, AgentProvider] = {} self._default_provider_type: ProviderType = config.workflow.runtime.provider.name self._resume_session_ids: dict[str, str] = {} + self._resume_session_cwds: dict[str, str] = {} @property def default_provider_type(self) -> ProviderType: @@ -129,6 +130,8 @@ async def _get_or_create_provider(self, provider_type: ProviderType) -> AgentPro # Pass stored resume session IDs to newly created providers if self._resume_session_ids and hasattr(provider, "set_resume_session_ids"): provider.set_resume_session_ids(self._resume_session_ids) # type: ignore[union-attr] + if self._resume_session_cwds and hasattr(provider, "set_resume_session_cwds"): + provider.set_resume_session_cwds(self._resume_session_cwds) # type: ignore[union-attr] self._providers[provider_type] = provider return provider @@ -194,3 +197,19 @@ def set_resume_session_ids(self, ids: dict[str, str]) -> None: for provider in self._providers.values(): if hasattr(provider, "set_resume_session_ids"): provider.set_resume_session_ids(ids) # type: ignore[union-attr] + + def set_resume_session_cwds(self, cwds: dict[str, str]) -> None: + """Store session working directories restored from a checkpoint. + + Forwarded to providers that support ``set_resume_session_cwds`` — + both already-active providers and providers created lazily in the + future. The mapping lets a provider skip resuming a session whose + working directory no longer matches the agent's resolved cwd. + + Args: + cwds: Mapping of agent names to session working directories. + """ + self._resume_session_cwds = dict(cwds) + for provider in self._providers.values(): + if hasattr(provider, "set_resume_session_cwds"): + provider.set_resume_session_cwds(cwds) # type: ignore[union-attr] diff --git a/tests/test_config/test_schema.py b/tests/test_config/test_schema.py index 03ce84be..307237d1 100644 --- a/tests/test_config/test_schema.py +++ b/tests/test_config/test_schema.py @@ -439,6 +439,16 @@ def test_custom_provider(self) -> None: assert config.provider.name == "openai-agents" assert config.default_model == "gpt-4" + def test_working_dir_defaults_to_none(self) -> None: + """Requirement: runtime.working_dir is an optional workflow-wide default.""" + assert RuntimeConfig().working_dir is None + + def test_working_dir_accepts_string(self) -> None: + """Requirement: runtime.working_dir accepts plain and templated paths.""" + assert RuntimeConfig(working_dir="/repo").working_dir == "/repo" + templated = RuntimeConfig(working_dir="{{ workflow.input.repo }}") + assert templated.working_dir == "{{ workflow.input.repo }}" + def test_invalid_provider_raises(self) -> None: """Test that invalid provider raises ValidationError.""" with pytest.raises(ValidationError): diff --git a/tests/test_config/test_validator_capabilities.py b/tests/test_config/test_validator_capabilities.py index afed41c3..c84b1883 100644 --- a/tests/test_config/test_validator_capabilities.py +++ b/tests/test_config/test_validator_capabilities.py @@ -1176,3 +1176,126 @@ def test_inline_non_llm_human_gate_skipped(self, patch_caps: Any) -> None: ), ) validate_workflow_config(config) # must not raise — human_gate is skipped + + +class TestWorkingDirCrossCheck: + """Requirement: ``working_dir`` (per-agent or runtime-wide) against a provider + declaring ``working_dir=False`` is a hard validate error — the setting would + otherwise be silently dropped (agent-mcp-working-dir, todo 1).""" + + def test_agent_working_dir_against_unsupported_provider_errors(self, patch_caps: Any) -> None: + patch_caps({"copilot": _caps(working_dir=False)}) + config = _build_workflow( + agents=[AgentDef(name="a", prompt="hi", working_dir="/repo")], + ) + with pytest.raises(ConfigurationError, match="working_dir"): + validate_workflow_config(config) + + def test_agent_working_dir_against_supported_provider_passes(self, patch_caps: Any) -> None: + patch_caps({"copilot": _caps(working_dir=True)}) + config = _build_workflow( + agents=[AgentDef(name="a", prompt="hi", working_dir="/repo")], + ) + validate_workflow_config(config) # no raise + + def test_runtime_working_dir_against_unsupported_provider_errors(self, patch_caps: Any) -> None: + """runtime.working_dir is inherited by every LLM agent on that provider.""" + patch_caps({"copilot": _caps(working_dir=False)}) + config = WorkflowConfig( + workflow=WorkflowDef( + name="test", + entry_point="a", + runtime=RuntimeConfig(provider="copilot", working_dir="/repo"), + ), + agents=[AgentDef(name="a", prompt="hi")], + ) + with pytest.raises(ConfigurationError, match="runtime.working_dir"): + validate_workflow_config(config) + + def test_runtime_working_dir_all_agents_override_to_capable_provider_passes( + self, patch_caps: Any + ) -> None: + """Default provider incapable but every LLM agent overrides to a capable + one → runtime.working_dir never reaches the incapable provider.""" + patch_caps( + { + "copilot": _caps(working_dir=False), + "claude": _caps(working_dir=True), + } + ) + config = WorkflowConfig( + workflow=WorkflowDef( + name="test", + entry_point="a", + runtime=RuntimeConfig(provider="copilot", working_dir="/repo"), + ), + agents=[AgentDef(name="a", prompt="hi", provider="claude")], + ) + validate_workflow_config(config) # no raise + + def test_per_agent_provider_override_against_working_dir_errors(self, patch_caps: Any) -> None: + """Agent overriding to a working_dir=False provider errors even when the + default provider supports working_dir.""" + patch_caps( + { + "copilot": _caps(working_dir=True), + "hermes": _caps(working_dir=False), + } + ) + config = _build_workflow( + agents=[AgentDef(name="a", prompt="hi", provider="hermes", working_dir="/repo")], + ) + with pytest.raises(ConfigurationError, match="hermes"): + validate_workflow_config(config) + + def test_for_each_inline_working_dir_against_unsupported_provider_errors( + self, patch_caps: Any + ) -> None: + """for_each inline agents get the same working_dir cross-check (#270 parity).""" + patch_caps({"copilot": _caps(working_dir=False)}) + config = _for_each_workflow( + inline=AgentDef(name="inner", prompt="{{ item }}", working_dir="/repo"), + ) + with pytest.raises(ConfigurationError, match="inner.*working_dir"): + validate_workflow_config(config) + + def test_for_each_inline_inherits_runtime_working_dir_errors(self, patch_caps: Any) -> None: + """An inline agent without its own working_dir still inherits the + runtime-wide default → error on a working_dir=False provider.""" + patch_caps({"copilot": _caps(working_dir=False), "claude": _caps(working_dir=True)}) + config = WorkflowConfig( + workflow=WorkflowDef( + name="test", + entry_point="entry", + runtime=RuntimeConfig(provider="copilot", working_dir="/repo"), + ), + agents=[AgentDef(name="entry", prompt="hi", provider="claude")], + for_each=[ + ForEachDef( + name="loop", + type="for_each", + source="entry.output.items", + **{"as": "item"}, + agent=AgentDef(name="inner", prompt="{{ item }}"), + ) + ], + ) + with pytest.raises(ConfigurationError, match="runtime.working_dir.*'inner'"): + validate_workflow_config(config) + + def test_no_working_dir_anywhere_against_unsupported_provider_passes( + self, patch_caps: Any + ) -> None: + """working_dir=False capability alone never errors without the setting.""" + patch_caps({"copilot": _caps(working_dir=False)}) + config = _build_workflow(agents=[AgentDef(name="a", prompt="hi")]) + validate_workflow_config(config) # no raise + + def test_script_agent_working_dir_skipped(self, patch_caps: Any) -> None: + """Script steps run a local subprocess (not a provider session) — the + capability gate must not fire for them even with working_dir set.""" + patch_caps({"copilot": _caps(working_dir=False)}) + config = _build_workflow( + agents=[AgentDef(name="s", type="script", command="ls", working_dir="/tmp")], + ) + validate_workflow_config(config) # no raise diff --git a/tests/test_config/test_wait_schema.py b/tests/test_config/test_wait_schema.py index 400f037b..d5d7a173 100644 --- a/tests/test_config/test_wait_schema.py +++ b/tests/test_config/test_wait_schema.py @@ -254,3 +254,55 @@ def test_minimal_wait_workflow(self) -> None: config = _make_workflow(wait) # Should not raise. validate_workflow_config(config) + + +class TestWorkingDirTypeMatrix: + """Requirement: ``working_dir`` is allowed on provider-backed LLM agents and + script steps, and rejected on wait/set/terminate/human_gate/workflow types.""" + + @pytest.mark.parametrize( + "kwargs", + [ + {"name": "llm", "prompt": "hi"}, + {"name": "script", "type": "script", "command": "ls"}, + ], + ids=["llm_agent", "script_step"], + ) + def test_working_dir_allowed(self, kwargs: dict) -> None: + agent = AgentDef(**kwargs, working_dir="/repo") # type: ignore[arg-type] + assert agent.working_dir == "/repo" + + @pytest.mark.parametrize( + "kwargs,match", + [ + ( + {"name": "w", "type": "wait", "duration": "1s"}, + "wait agents cannot have 'working_dir'", + ), + ( + {"name": "s", "type": "set", "value": "1"}, + "set agents cannot have 'working_dir'", + ), + ( + {"name": "t", "type": "terminate", "status": "success", "reason": "done"}, + "terminate agents cannot have 'working_dir'", + ), + ( + { + "name": "g", + "type": "human_gate", + "prompt": "Pick", + "options": [GateOption(label="Yes", value="yes", route="$end")], + }, + "human_gate agents cannot have 'working_dir'", + ), + ( + {"name": "wf", "type": "workflow", "workflow": "./sub.yaml"}, + "workflow agents cannot have 'working_dir'", + ), + ], + ids=["wait", "set", "terminate", "human_gate", "workflow"], + ) + def test_working_dir_rejected(self, kwargs: dict, match: str) -> None: + with pytest.raises(PydanticValidationError, match=match): + AgentDef(**kwargs, working_dir="/repo") # type: ignore[arg-type] diff --git a/tests/test_engine/test_agent_timeout.py b/tests/test_engine/test_agent_timeout.py index e9d86cc4..baadf94d 100644 --- a/tests/test_engine/test_agent_timeout.py +++ b/tests/test_engine/test_agent_timeout.py @@ -644,7 +644,9 @@ def mock_handler(agent, prompt, context): async def patched_get_executor(agent): executor = await original_get_executor(agent) - if agent.name == "processor": + # for_each qualifies the per-iteration name ("processor[0]") before + # resolving the executor, so match the base agent name by prefix. + if agent.name.startswith("processor"): original_exec = executor.execute async def slow_exec(*args, **kwargs): @@ -721,7 +723,9 @@ def mock_handler(agent, prompt, context): async def patched_get_executor(agent): executor = await original_get_executor(agent) - if agent.name == "processor": + # for_each qualifies the per-iteration name ("processor[0]") before + # resolving the executor, so match the base agent name by prefix. + if agent.name.startswith("processor"): original_exec = executor.execute async def slow_exec(*args, **kwargs): diff --git a/tests/test_engine/test_checkpoint.py b/tests/test_engine/test_checkpoint.py index 3519e292..2c8f27d7 100644 --- a/tests/test_engine/test_checkpoint.py +++ b/tests/test_engine/test_checkpoint.py @@ -275,6 +275,49 @@ def test_copilot_session_ids_included(self, tmp_path: Path) -> None: data = json.loads(path.read_text()) assert data["copilot_session_ids"] == {"agent_a": "sid-123"} + def test_copilot_session_cwds_round_trip(self, tmp_path: Path) -> None: + """Requirement: copilot_session_cwds persists through save+load (Persistent field).""" + wf = _write_workflow(tmp_path) + ctx = _make_context() + limits = _make_limits() + error = RuntimeError("err") + + with patch.object(CheckpointManager, "get_checkpoints_dir", return_value=tmp_path): + path = CheckpointManager.save_checkpoint( + wf, + ctx, + limits, + "a", + error, + {}, + copilot_session_ids={"agent_a": "sid-123"}, + copilot_session_cwds={"agent_a": "/repo/a"}, + ) + + assert path is not None + data = json.loads(path.read_text()) + assert data["copilot_session_cwds"] == {"agent_a": "/repo/a"} + + cp = CheckpointManager.load_checkpoint(path) + assert cp.copilot_session_cwds == {"agent_a": "/repo/a"} + + def test_copilot_session_cwds_defaults_to_empty(self, tmp_path: Path) -> None: + """Requirement: checkpoints saved without cwds load with an empty mapping.""" + wf = _write_workflow(tmp_path) + ctx = _make_context() + limits = _make_limits() + error = RuntimeError("err") + + with patch.object(CheckpointManager, "get_checkpoints_dir", return_value=tmp_path): + path = CheckpointManager.save_checkpoint(wf, ctx, limits, "a", error, {}) + + assert path is not None + data = json.loads(path.read_text()) + assert data["copilot_session_cwds"] == {} + + cp = CheckpointManager.load_checkpoint(path) + assert cp.copilot_session_cwds == {} + def test_run_id_and_event_log_path_included(self, tmp_path: Path) -> None: """``run_id`` and ``event_log_path`` round-trip through save+load. @@ -427,6 +470,36 @@ def test_loads_legacy_checkpoint_without_run_id_or_event_log_path(self, tmp_path assert cp.run_id == "" assert cp.event_log_path == "" + def test_loads_legacy_checkpoint_without_copilot_session_cwds(self, tmp_path: Path) -> None: + """Requirement: pre-cwd checkpoints (no copilot_session_cwds key) load with default. + + Backward compatibility: ``load_checkpoint`` must default the mapping + to ``{}`` without raising so resume falls back to the legacy + resume-by-session-id behavior. + """ + legacy = tmp_path / "legacy-no-cwds.json" + legacy.write_text( + json.dumps( + { + "version": 1, + "workflow_path": "/x.yaml", + "workflow_hash": "sha256:abc", + "created_at": "2026-01-01T00:00:00+00:00", + "failure": {"error_type": "X", "message": "m", "agent": "a", "iteration": 0}, + "inputs": {}, + "current_agent": "a", + "context": {"workflow_inputs": {}, "agent_outputs": {}}, + "limits": {"current_iteration": 0, "max_iterations": 10}, + "copilot_session_ids": {"a": "sid-1"}, + # No copilot_session_cwds — pre-cwd shape. + } + ) + ) + + cp = CheckpointManager.load_checkpoint(legacy) + assert cp.copilot_session_ids == {"a": "sid-1"} + assert cp.copilot_session_cwds == {} + # --------------------------------------------------------------------------- # CheckpointManager.find_latest_checkpoint tests diff --git a/tests/test_engine/test_event_log.py b/tests/test_engine/test_event_log.py index f66eb544..035c7617 100644 --- a/tests/test_engine/test_event_log.py +++ b/tests/test_engine/test_event_log.py @@ -266,3 +266,47 @@ def test_falls_back_when_no_existing_run_id(self, tmp_path, monkeypatch): assert sub.path.exists() finally: sub.close() + + def test_records_working_dir_on_agent_start_events(self, tmp_path, monkeypatch): + """Requirement: the JSONL log preserves the resolved ``working_dir`` + carried by the additive ``parallel_agent_started`` / + ``for_each_agent_started`` events, including an explicit ``None`` when + no working directory was configured.""" + monkeypatch.setenv("TMPDIR", str(tmp_path)) + sub = EventLogSubscriber("working-dir-events") + + sub.on_event( + WorkflowEvent( + type="parallel_agent_started", + timestamp=time.time(), + data={ + "group_name": "fan", + "agent_name": "explicit_a", + "working_dir": "/repo/a", + }, + ) + ) + sub.on_event( + WorkflowEvent( + type="for_each_agent_started", + timestamp=time.time(), + data={ + "group_name": "fans", + "agent_name": "fan_agent[0]", + "item_key": "0", + "working_dir": None, + }, + ) + ) + sub.close() + + lines = sub.path.read_text().strip().split("\n") + assert len(lines) == 2 + parallel = json.loads(lines[0]) + assert parallel["type"] == "parallel_agent_started" + assert parallel["data"]["working_dir"] == "/repo/a" + for_each = json.loads(lines[1]) + assert for_each["type"] == "for_each_agent_started" + assert for_each["data"]["agent_name"] == "fan_agent[0]" + assert for_each["data"]["item_key"] == "0" + assert for_each["data"]["working_dir"] is None diff --git a/tests/test_engine/test_subworkflow.py b/tests/test_engine/test_subworkflow.py index 8eaed90c..73ac0488 100644 --- a/tests/test_engine/test_subworkflow.py +++ b/tests/test_engine/test_subworkflow.py @@ -2510,3 +2510,80 @@ async def test_failed_terminate_in_for_each_workflow_iteration( assert "iteration x failed" in message or "iteration y failed" in message, ( f"expected child reason in error chain; got: {message}" ) + + +class TestSubWorkflowWorkingDir: + @pytest.mark.asyncio + async def test_subworkflow_no_working_dir_inheritance(self, tmp_workflow_dir: Path) -> None: + """Requirement: a sub-workflow does not inherit the parent's runtime.working_dir; + + the child's own relative working_dir resolves against the child workflow file's directory. + """ + child_dir = tmp_workflow_dir / "child_subdir" + child_dir.mkdir() + + target_dir = child_dir / "x" + target_dir.mkdir() + + _write_yaml( + child_dir / "child.yaml", + """\ + workflow: + name: child-workflow + entry_point: child_agent + runtime: + provider: copilot + working_dir: "./x" + limits: + max_iterations: 5 + agents: + - name: child_agent + prompt: "Child prompt" + routes: + - to: "$end" + output: + result: "{{ child_agent.output.result }}" + """, + ) + + parent_path = tmp_workflow_dir / "parent.yaml" + parent_path.write_text("dummy", encoding="utf-8") + + config = WorkflowConfig( + workflow=WorkflowDef( + name="parent-workflow", + entry_point="sub_wf", + runtime=RuntimeConfig( + provider="copilot", + working_dir="/some/parent/dir", + ), + context=ContextConfig(mode="accumulate"), + limits=LimitsConfig(max_iterations=10), + ), + agents=[ + AgentDef( + name="sub_wf", + type="workflow", + workflow="child_subdir/child.yaml", + routes=[RouteDef(to="$end")], + ), + ], + output={ + "result": "{{ sub_wf.output.result }}", + }, + ) + + resolved_cwds = [] + + def mock_handler(agent, prompt, context): + resolved_cwds.append(agent.working_dir) + return {"result": "ok"} + + provider = CopilotProvider(mock_handler=mock_handler) + engine = WorkflowEngine(config, provider, workflow_path=parent_path) + + result = await engine.run({}) + + assert result["result"] == "ok" + assert len(resolved_cwds) == 1 + assert resolved_cwds[0] == str(target_dir.resolve()) diff --git a/tests/test_engine/test_validator_integration.py b/tests/test_engine/test_validator_integration.py index 0236c6e0..b35976e5 100644 --- a/tests/test_engine/test_validator_integration.py +++ b/tests/test_engine/test_validator_integration.py @@ -524,3 +524,89 @@ async def exec_fn(*, agent: AgentDef, rendered_prompt: str, **kw: Any) -> AgentO # Grader timed out → fail open (treated as pass) → original returned, no re-run. assert result is original + + +class TestValidatorWorkingDirectoryInheritance: + """Verify that the synthetic validator agent inherits the primary agent's working directory.""" + + @pytest.mark.asyncio + async def test_validator_working_dir_inheritance(self) -> None: + # Requirement: The synthetic validator agent must inherit the working directory of + # the primary agent. + captured_validator_agent: AgentDef | None = None + + async def exec_fn(*, agent: AgentDef, rendered_prompt: str, **kw: Any) -> AgentOutput: + nonlocal captured_validator_agent + captured_validator_agent = agent + return AgentOutput( + content={"passed": True, "issues": []}, + raw_response="", + model="judge", + ) + + primary_agent = AgentDef( + name="reviewer", + model="gpt-4", + prompt="Review", + output={"summary": OutputField(type="string")}, + validator=ValidatorConfig(criteria="check"), + routes=[RouteDef(to="$end")], + working_dir="/path/to/resolved/working_dir", + ) + provider = CopilotProvider(mock_handler=lambda a, p, c: {}) + provider.execute = exec_fn # type: ignore[method-assign] + + from conductor.engine.validator import OutputValidator + + validator = OutputValidator() + + await validator.validate( + agent=primary_agent, + primary_prompt="Review", + primary_output={"summary": "orig"}, + provider=provider, + ) + + assert captured_validator_agent is not None + assert captured_validator_agent.working_dir == "/path/to/resolved/working_dir" + + @pytest.mark.asyncio + async def test_validator_working_dir_none(self) -> None: + # Requirement: If the primary agent has no working directory set, the validator's + # working directory is None. + captured_validator_agent: AgentDef | None = None + + async def exec_fn(*, agent: AgentDef, rendered_prompt: str, **kw: Any) -> AgentOutput: + nonlocal captured_validator_agent + captured_validator_agent = agent + return AgentOutput( + content={"passed": True, "issues": []}, + raw_response="", + model="judge", + ) + + primary_agent = AgentDef( + name="reviewer", + model="gpt-4", + prompt="Review", + output={"summary": OutputField(type="string")}, + validator=ValidatorConfig(criteria="check"), + routes=[RouteDef(to="$end")], + working_dir=None, + ) + provider = CopilotProvider(mock_handler=lambda a, p, c: {}) + provider.execute = exec_fn # type: ignore[method-assign] + + from conductor.engine.validator import OutputValidator + + validator = OutputValidator() + + await validator.validate( + agent=primary_agent, + primary_prompt="Review", + primary_output={"summary": "orig"}, + provider=provider, + ) + + assert captured_validator_agent is not None + assert captured_validator_agent.working_dir is None diff --git a/tests/test_engine/test_workflow.py b/tests/test_engine/test_workflow.py index 62ef35f8..2b7b836f 100644 --- a/tests/test_engine/test_workflow.py +++ b/tests/test_engine/test_workflow.py @@ -9,13 +9,16 @@ """ import asyncio +import os import sys +from pathlib import Path import pytest from conductor.config.schema import ( AgentDef, ContextConfig, + ForEachDef, GateOption, HooksConfig, InputDef, @@ -28,7 +31,9 @@ WorkflowDef, ) from conductor.engine.workflow import WorkflowEngine +from conductor.events import WorkflowEventEmitter from conductor.exceptions import ExecutionError +from conductor.providers.base import AgentOutput from conductor.providers.copilot import CopilotProvider @@ -3185,7 +3190,7 @@ async def test_failed_terminate_does_not_save_checkpoint(self, tmp_path) -> None config = self._config_with_terminate("failed", reason="halt") provider = CopilotProvider(mock_handler=lambda *_a, **_kw: {"value": "v"}) - wf_file = tmp_path / "wf.yaml" + wf_file = tmp_path / "failed_terminate_checkpoint_test.yaml" wf_file.write_text("name: t\n") engine = WorkflowEngine(config, provider, workflow_path=wf_file) @@ -3610,3 +3615,807 @@ async def test_lifecycle_event_ordering_failed_terminate(self) -> None: assert af_index < wf_index, ( f"agent_failed must precede workflow_failed; got order: {types_in_order}" ) + + +class _RecordingWorkingDirProvider: + """Minimal provider that records the ``working_dir`` of every agent passed to it. + + Duck-types the ``AgentProvider`` contract the engine consumes. Returns one + structured field per declared output key so ``engine.run({})`` reaches a + clean ``$end``. + """ + + def __init__(self) -> None: + self.seen: list[tuple[str, str | None]] = [] + self.calls: int = 0 + + async def execute( + self, + agent, + context, + rendered_prompt, + tools=None, + interrupt_signal=None, + event_callback=None, + ): + self.calls += 1 + self.seen.append((agent.name, agent.working_dir)) + content = dict.fromkeys(agent.output or {}, f"{agent.name}-ok") + return AgentOutput( + content=content, + raw_response=None, + model=agent.model, + input_tokens=1, + output_tokens=1, + ) + + async def validate_connection(self) -> bool: + return True + + async def close(self) -> None: + return None + + async def get_max_prompt_tokens(self, model: str): + return None + + +def _single_agent_config( + *, + working_dir: str | None = None, + runtime_working_dir: str | None = None, + model: str = "gpt-4", + max_tokens: int | None = None, +) -> WorkflowConfig: + """One LLM agent routing straight to ``$end``.""" + return WorkflowConfig( + workflow=WorkflowDef( + name="wd-single", + entry_point="worker", + runtime=RuntimeConfig(provider="copilot", working_dir=runtime_working_dir), + context=ContextConfig(mode="accumulate", max_tokens=max_tokens), + limits=LimitsConfig(max_iterations=10), + ), + agents=[ + AgentDef( + name="worker", + model=model, + prompt="Do work", + working_dir=working_dir, + output={"result": OutputField(type="string")}, + routes=[RouteDef(to="$end")], + ), + ], + output={"result": "{{ worker.output.result }}"}, + ) + + +def _workflow_file(tmp_path: Path) -> Path: + wf_file = tmp_path / "workflow.yaml" + wf_file.write_text("name: wd\n") + return wf_file + + +class TestAgentWorkingDirResolution: + """Engine resolution of ``AgentDef.working_dir`` / ``runtime.working_dir``. + + Covers linear, parallel, and for-each execution paths, the + ``agent > runtime`` precedence, Jinja rendering against the per-agent + context, relative-path resolution against the workflow file's directory, + ``~`` expansion, and the missing-directory ``ExecutionError``. + """ + + @pytest.mark.asyncio + async def test_linear_absolute_working_dir_reaches_provider(self, tmp_path: Path) -> None: + """Requirement: an absolute ``working_dir`` on the agent is resolved by the + engine and reaches the provider (the resolved value is set on AgentDef).""" + target = tmp_path / "repo" + target.mkdir() + provider = _RecordingWorkingDirProvider() + engine = WorkflowEngine( + _single_agent_config(working_dir=str(target)), + provider, + workflow_path=_workflow_file(tmp_path), + ) + + await engine.run({}) + + assert provider.seen == [("worker", os.path.normpath(str(target)))] + + @pytest.mark.asyncio + async def test_agent_beats_runtime_precedence(self, tmp_path: Path) -> None: + """Requirement: precedence ``agent.working_dir`` > ``runtime.working_dir`` + — when both are set, the provider sees the agent's value.""" + agent_dir = tmp_path / "agent-dir" + runtime_dir = tmp_path / "runtime-dir" + agent_dir.mkdir() + runtime_dir.mkdir() + provider = _RecordingWorkingDirProvider() + engine = WorkflowEngine( + _single_agent_config(working_dir=str(agent_dir), runtime_working_dir=str(runtime_dir)), + provider, + workflow_path=_workflow_file(tmp_path), + ) + + await engine.run({}) + + assert provider.seen == [("worker", os.path.normpath(str(agent_dir)))] + + @pytest.mark.asyncio + async def test_templated_runtime_working_dir_is_rendered(self, tmp_path: Path) -> None: + """Requirement (Oracle r2 amendment): ``runtime.working_dir`` is also + Jinja-rendered against the agent's context, not passed through raw.""" + target = tmp_path / "from-input" + target.mkdir() + provider = _RecordingWorkingDirProvider() + engine = WorkflowEngine( + _single_agent_config(runtime_working_dir="{{ workflow.input.target }}"), + provider, + workflow_path=_workflow_file(tmp_path), + ) + + await engine.run({"target": str(target)}) + + assert provider.seen == [("worker", os.path.normpath(str(target)))] + + @pytest.mark.asyncio + async def test_relative_working_dir_resolves_against_workflow_dir(self, tmp_path: Path) -> None: + """Requirement: a relative ``working_dir`` resolves against the workflow + file's directory (``self._workflow_dir``), not the process's current cwd.""" + (tmp_path / "sub").mkdir() + provider = _RecordingWorkingDirProvider() + engine = WorkflowEngine( + _single_agent_config(working_dir="./sub"), + provider, + workflow_path=_workflow_file(tmp_path), + ) + + await engine.run({}) + + assert provider.seen == [("worker", os.path.normpath(str(tmp_path / "sub")))] + + @pytest.mark.asyncio + async def test_expanduser_tilde(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + """Requirement: ``~`` in ``working_dir`` is expanded via + ``Path.expanduser()`` to an absolute home-directory path.""" + fake_home = tmp_path / "home" + (fake_home / "proj").mkdir(parents=True) + monkeypatch.setenv("HOME", str(fake_home)) + provider = _RecordingWorkingDirProvider() + engine = WorkflowEngine( + _single_agent_config(working_dir="~/proj"), + provider, + workflow_path=_workflow_file(tmp_path), + ) + + await engine.run({}) + + assert provider.seen == [("worker", os.path.normpath(str(fake_home / "proj")))] + + @pytest.mark.asyncio + async def test_missing_dir_raises_before_provider_call(self, tmp_path: Path) -> None: + """Requirement: a nonexistent directory raises ``ExecutionError`` BEFORE + the provider call (the provider must not be called at all).""" + provider = _RecordingWorkingDirProvider() + engine = WorkflowEngine( + _single_agent_config(working_dir=str(tmp_path / "does-not-exist")), + provider, + workflow_path=_workflow_file(tmp_path), + ) + + with pytest.raises(ExecutionError, match="working_dir"): + await engine.run({}) + + assert provider.calls == 0 + + @pytest.mark.asyncio + async def test_no_working_dir_anywhere_passes_none(self, tmp_path: Path) -> None: + """Requirement: when neither agent nor runtime sets ``working_dir``, the + provider receives ``None`` (and falls back to ``os.getcwd()``).""" + provider = _RecordingWorkingDirProvider() + engine = WorkflowEngine( + _single_agent_config(), + provider, + workflow_path=_workflow_file(tmp_path), + ) + + await engine.run({}) + + assert provider.seen == [("worker", None)] + + @pytest.mark.asyncio + async def test_parallel_agents_resolve_their_own_working_dir(self, tmp_path: Path) -> None: + """Requirement: in a parallel group each agent is resolved individually + after its own per-agent context is built (agent-level > runtime-level).""" + dir_a = tmp_path / "a" + dir_b = tmp_path / "b" + runtime_dir = tmp_path / "runtime" + dir_a.mkdir() + dir_b.mkdir() + runtime_dir.mkdir() + config = WorkflowConfig( + workflow=WorkflowDef( + name="wd-parallel", + entry_point="fan", + runtime=RuntimeConfig(provider="copilot", working_dir=str(runtime_dir)), + context=ContextConfig(mode="accumulate"), + limits=LimitsConfig(max_iterations=10), + ), + agents=[ + AgentDef( + name="explicit_a", + model="gpt-4", + prompt="A", + working_dir=str(dir_a), + output={"r": OutputField(type="string")}, + ), + AgentDef( + name="inherit_b", + model="gpt-4", + prompt="B", + output={"r": OutputField(type="string")}, + ), + ], + parallel=[ + ParallelGroup( + name="fan", + agents=["explicit_a", "inherit_b"], + routes=[RouteDef(to="$end")], + ), + ], + output={}, + ) + provider = _RecordingWorkingDirProvider() + engine = WorkflowEngine(config, provider, workflow_path=_workflow_file(tmp_path)) + + await engine.run({}) + + assert sorted(provider.seen) == [ + ("explicit_a", os.path.normpath(str(dir_a))), + ("inherit_b", os.path.normpath(str(runtime_dir))), + ] + + @pytest.mark.asyncio + async def test_for_each_resolves_item_template_per_iteration(self, tmp_path: Path) -> None: + """Requirement: in for_each, ``{{ item }}`` in the path resolves AFTER the + loop variables are substituted — each iteration gets its own directory.""" + for name in ("one", "two"): + (tmp_path / name).mkdir() + config = WorkflowConfig( + workflow=WorkflowDef( + name="wd-for-each", + entry_point="lister", + runtime=RuntimeConfig(provider="copilot"), + context=ContextConfig(mode="accumulate"), + limits=LimitsConfig(max_iterations=10), + ), + agents=[ + AgentDef( + name="lister", + model="gpt-4", + prompt="List", + output={"repos": OutputField(type="array")}, + routes=[RouteDef(to="fans")], + ), + ], + for_each=[ + ForEachDef( + name="fans", + type="for_each", + source="lister.output.repos", + **{"as": "repo"}, + agent=AgentDef( + name="fan_agent", + model="gpt-4", + prompt="Work {{ repo }}", + working_dir=str(tmp_path / "{{ repo }}"), + output={"r": OutputField(type="string")}, + ), + routes=[RouteDef(to="$end")], + ), + ], + output={}, + ) + provider = _RecordingWorkingDirProvider() + + async def _execute(agent, context, rendered_prompt, tools=None, **kwargs): + if agent.name == "lister": + return AgentOutput( + content={"repos": ["one", "two"]}, + raw_response=None, + model=agent.model, + input_tokens=1, + output_tokens=1, + ) + provider.seen.append((agent.name, agent.working_dir)) + return AgentOutput( + content={"r": "ok"}, + raw_response=None, + model=agent.model, + input_tokens=1, + output_tokens=1, + ) + + provider.execute = _execute # type: ignore[method-assign] + engine = WorkflowEngine(config, provider, workflow_path=_workflow_file(tmp_path)) + + await engine.run({}) + + assert sorted(provider.seen) == [ + ("fan_agent[0]", os.path.normpath(str(tmp_path / "one"))), + ("fan_agent[1]", os.path.normpath(str(tmp_path / "two"))), + ] + + @pytest.mark.asyncio + async def test_templated_model_does_not_clobber_resolved_working_dir( + self, tmp_path: Path + ) -> None: + """Requirement (Oracle r3 amendment, regression): a templated ``model`` is + rendered in ``AgentExecutor`` via ``model_copy(update={...})`` — a merge, + so the engine-resolved ``working_dir`` must survive to the provider.""" + target = tmp_path / "repo" + target.mkdir() + provider = _RecordingWorkingDirProvider() + engine = WorkflowEngine( + _single_agent_config(working_dir=str(target), model="{{ workflow.input.model }}"), + provider, + workflow_path=_workflow_file(tmp_path), + ) + + await engine.run({"model": "gpt-4o"}) + + assert provider.seen == [("worker", os.path.normpath(str(target)))] + + @pytest.mark.asyncio + async def test_agent_started_carries_working_dir_and_context_window_after_trim( + self, tmp_path: Path + ) -> None: + """Requirement (Oracle r1 amendment, trim ordering): with a small + ``max_tokens``, the ``agent_started`` event carries BOTH the resolved + ``working_dir`` AND ``context_window_max``, and the context trim ran + BEFORE agent_context was built (the prompt is already trimmed).""" + target = tmp_path / "repo" + target.mkdir() + config = WorkflowConfig( + workflow=WorkflowDef( + name="wd-trim", + entry_point="first", + runtime=RuntimeConfig(provider="copilot"), + context=ContextConfig(mode="accumulate", max_tokens=1), + limits=LimitsConfig(max_iterations=10), + ), + agents=[ + AgentDef( + name="first", + model="gpt-4", + prompt="First", + output={"blob": OutputField(type="string")}, + routes=[RouteDef(to="second")], + ), + AgentDef( + name="second", + model="gpt-4", + prompt="Second", + working_dir=str(target), + output={"result": OutputField(type="string")}, + routes=[RouteDef(to="$end")], + ), + ], + output={"result": "{{ second.output.result }}"}, + ) + provider = _RecordingWorkingDirProvider() + events: list = [] + emitter = WorkflowEventEmitter() + emitter.subscribe(events.append) + engine = WorkflowEngine( + config, provider, event_emitter=emitter, workflow_path=_workflow_file(tmp_path) + ) + + await engine.run({}) + + started = [ + e for e in events if e.type == "agent_started" and e.data["agent_name"] == "second" + ] + assert len(started) == 1 + assert started[0].data["working_dir"] == os.path.normpath(str(target)) + assert "context_window_max" in started[0].data + assert provider.seen[-1][0] == "second" + + +class TestWorkingDirEvents: + """Observability events carrying the engine-resolved ``working_dir``. + + The linear path is covered by ``agent_started`` (todo 2, commit 075436b); + this class pins the two additive LLM-only events — + ``parallel_agent_started`` and ``for_each_agent_started`` — emitted right + after each per-agent/per-item ``working_dir`` resolution, plus regression + guards that the pre-existing envelope events keep their exact payloads. + """ + + @pytest.mark.asyncio + async def test_parallel_agent_started_carries_resolved_working_dir( + self, tmp_path: Path + ) -> None: + """Requirement: ``parallel_agent_started`` fires per LLM member AFTER its + own resolution and carries ``group_name``, ``agent_name`` and the + resolved ``working_dir`` (agent-level beats the runtime default).""" + dir_a = tmp_path / "a" + runtime_dir = tmp_path / "runtime" + dir_a.mkdir() + runtime_dir.mkdir() + config = WorkflowConfig( + workflow=WorkflowDef( + name="wd-ev-parallel", + entry_point="fan", + runtime=RuntimeConfig(provider="copilot", working_dir=str(runtime_dir)), + context=ContextConfig(mode="accumulate"), + limits=LimitsConfig(max_iterations=10), + ), + agents=[ + AgentDef( + name="explicit_a", + model="gpt-4", + prompt="A", + working_dir=str(dir_a), + output={"r": OutputField(type="string")}, + ), + AgentDef( + name="inherit_b", + model="gpt-4", + prompt="B", + output={"r": OutputField(type="string")}, + ), + ], + parallel=[ + ParallelGroup( + name="fan", + agents=["explicit_a", "inherit_b"], + routes=[RouteDef(to="$end")], + ), + ], + output={}, + ) + provider = _RecordingWorkingDirProvider() + events: list = [] + emitter = WorkflowEventEmitter() + emitter.subscribe(events.append) + engine = WorkflowEngine( + config, provider, event_emitter=emitter, workflow_path=_workflow_file(tmp_path) + ) + + await engine.run({}) + + started = sorted( + (e for e in events if e.type == "parallel_agent_started"), + key=lambda e: e.data["agent_name"], + ) + observed = [ + (e.data["group_name"], e.data["agent_name"], e.data["working_dir"]) for e in started + ] + assert observed == [ + ("fan", "explicit_a", os.path.normpath(str(dir_a))), + ("fan", "inherit_b", os.path.normpath(str(runtime_dir))), + ] + + @pytest.mark.asyncio + async def test_parallel_agent_started_working_dir_none_when_unset(self, tmp_path: Path) -> None: + """Requirement: without any ``working_dir`` the event stays valid and + carries ``working_dir=None`` (the provider falls back to its own cwd).""" + config = WorkflowConfig( + workflow=WorkflowDef( + name="wd-ev-parallel-none", + entry_point="fan", + runtime=RuntimeConfig(provider="copilot"), + context=ContextConfig(mode="accumulate"), + limits=LimitsConfig(max_iterations=10), + ), + agents=[ + AgentDef( + name="solo", + model="gpt-4", + prompt="Solo", + output={"r": OutputField(type="string")}, + ), + AgentDef( + name="other", + model="gpt-4", + prompt="Other", + output={"r": OutputField(type="string")}, + ), + ], + parallel=[ + ParallelGroup(name="fan", agents=["solo", "other"], routes=[RouteDef(to="$end")]), + ], + output={}, + ) + provider = _RecordingWorkingDirProvider() + events: list = [] + emitter = WorkflowEventEmitter() + emitter.subscribe(events.append) + engine = WorkflowEngine( + config, provider, event_emitter=emitter, workflow_path=_workflow_file(tmp_path) + ) + + await engine.run({}) + + started = [e for e in events if e.type == "parallel_agent_started"] + assert len(started) == 2 + solo = next(e for e in started if e.data["agent_name"] == "solo") + assert solo.data["group_name"] == "fan" + assert solo.data["working_dir"] is None + + @pytest.mark.asyncio + async def test_parallel_started_payload_unchanged_regression(self, tmp_path: Path) -> None: + """Requirement (regression): the existing ``parallel_started`` envelope + event keeps its exact pre-change payload (``group_name`` + ``agents`` + only) — ``working_dir`` observability is strictly additive.""" + config = WorkflowConfig( + workflow=WorkflowDef( + name="wd-ev-parallel-reg", + entry_point="fan", + runtime=RuntimeConfig(provider="copilot"), + context=ContextConfig(mode="accumulate"), + limits=LimitsConfig(max_iterations=10), + ), + agents=[ + AgentDef( + name="solo", + model="gpt-4", + prompt="Solo", + output={"r": OutputField(type="string")}, + ), + AgentDef( + name="other", + model="gpt-4", + prompt="Other", + output={"r": OutputField(type="string")}, + ), + ], + parallel=[ + ParallelGroup(name="fan", agents=["solo", "other"], routes=[RouteDef(to="$end")]), + ], + output={}, + ) + provider = _RecordingWorkingDirProvider() + events: list = [] + emitter = WorkflowEventEmitter() + emitter.subscribe(events.append) + engine = WorkflowEngine( + config, provider, event_emitter=emitter, workflow_path=_workflow_file(tmp_path) + ) + + await engine.run({}) + + envelope = [e for e in events if e.type == "parallel_started"] + assert len(envelope) == 1 + assert envelope[0].data == {"group_name": "fan", "agents": ["solo", "other"]} + + @pytest.mark.asyncio + async def test_for_each_agent_started_carries_resolved_working_dir_per_item( + self, tmp_path: Path + ) -> None: + """Requirement: ``for_each_agent_started`` fires per item AFTER the + ``{{ item }}`` template resolves and carries ``group_name``, the + qualified ``agent_name`` (``name[key]``), ``item_key`` and the resolved + per-iteration ``working_dir``.""" + for name in ("one", "two"): + (tmp_path / name).mkdir() + config = WorkflowConfig( + workflow=WorkflowDef( + name="wd-ev-for-each", + entry_point="lister", + runtime=RuntimeConfig(provider="copilot"), + context=ContextConfig(mode="accumulate"), + limits=LimitsConfig(max_iterations=10), + ), + agents=[ + AgentDef( + name="lister", + model="gpt-4", + prompt="List", + output={"repos": OutputField(type="array")}, + routes=[RouteDef(to="fans")], + ), + ], + for_each=[ + ForEachDef( + name="fans", + type="for_each", + source="lister.output.repos", + **{"as": "repo"}, + agent=AgentDef( + name="fan_agent", + model="gpt-4", + prompt="Work {{ repo }}", + working_dir=str(tmp_path / "{{ repo }}"), + output={"r": OutputField(type="string")}, + ), + routes=[RouteDef(to="$end")], + ), + ], + output={}, + ) + provider = _RecordingWorkingDirProvider() + + async def _execute(agent, context, rendered_prompt, tools=None, **kwargs): + if agent.name == "lister": + return AgentOutput( + content={"repos": ["one", "two"]}, + raw_response=None, + model=agent.model, + input_tokens=1, + output_tokens=1, + ) + return AgentOutput( + content={"r": "ok"}, + raw_response=None, + model=agent.model, + input_tokens=1, + output_tokens=1, + ) + + provider.execute = _execute # type: ignore[method-assign] + events: list = [] + emitter = WorkflowEventEmitter() + emitter.subscribe(events.append) + engine = WorkflowEngine( + config, provider, event_emitter=emitter, workflow_path=_workflow_file(tmp_path) + ) + + await engine.run({}) + + started = sorted( + (e for e in events if e.type == "for_each_agent_started"), + key=lambda e: e.data["item_key"], + ) + assert [ + (e.data["group_name"], e.data["agent_name"], e.data["item_key"], e.data["working_dir"]) + for e in started + ] == [ + ("fans", "fan_agent[0]", "0", os.path.normpath(str(tmp_path / "one"))), + ("fans", "fan_agent[1]", "1", os.path.normpath(str(tmp_path / "two"))), + ] + + @pytest.mark.asyncio + async def test_for_each_agent_started_working_dir_none_when_unset(self, tmp_path: Path) -> None: + """Requirement: without any ``working_dir`` the per-item event stays + valid and carries ``working_dir=None``.""" + config = WorkflowConfig( + workflow=WorkflowDef( + name="wd-ev-for-each-none", + entry_point="lister", + runtime=RuntimeConfig(provider="copilot"), + context=ContextConfig(mode="accumulate"), + limits=LimitsConfig(max_iterations=10), + ), + agents=[ + AgentDef( + name="lister", + model="gpt-4", + prompt="List", + output={"repos": OutputField(type="array")}, + routes=[RouteDef(to="fans")], + ), + ], + for_each=[ + ForEachDef( + name="fans", + type="for_each", + source="lister.output.repos", + **{"as": "repo"}, + agent=AgentDef( + name="fan_agent", + model="gpt-4", + prompt="Work {{ repo }}", + output={"r": OutputField(type="string")}, + ), + routes=[RouteDef(to="$end")], + ), + ], + output={}, + ) + provider = _RecordingWorkingDirProvider() + + async def _execute(agent, context, rendered_prompt, tools=None, **kwargs): + if agent.name == "lister": + return AgentOutput( + content={"repos": ["one"]}, + raw_response=None, + model=agent.model, + input_tokens=1, + output_tokens=1, + ) + return AgentOutput( + content={"r": "ok"}, + raw_response=None, + model=agent.model, + input_tokens=1, + output_tokens=1, + ) + + provider.execute = _execute # type: ignore[method-assign] + events: list = [] + emitter = WorkflowEventEmitter() + emitter.subscribe(events.append) + engine = WorkflowEngine( + config, provider, event_emitter=emitter, workflow_path=_workflow_file(tmp_path) + ) + + await engine.run({}) + + started = [e for e in events if e.type == "for_each_agent_started"] + assert len(started) == 1 + assert started[0].data["agent_name"] == "fan_agent[0]" + assert started[0].data["item_key"] == "0" + assert started[0].data["working_dir"] is None + + @pytest.mark.asyncio + async def test_for_each_item_started_payload_unchanged_regression(self, tmp_path: Path) -> None: + """Requirement (regression): the existing ``for_each_item_started`` + envelope event keeps its exact pre-change payload (``group_name`` + + ``item_key`` + ``index``) — the new event is strictly additive.""" + config = WorkflowConfig( + workflow=WorkflowDef( + name="wd-ev-for-each-reg", + entry_point="lister", + runtime=RuntimeConfig(provider="copilot"), + context=ContextConfig(mode="accumulate"), + limits=LimitsConfig(max_iterations=10), + ), + agents=[ + AgentDef( + name="lister", + model="gpt-4", + prompt="List", + output={"repos": OutputField(type="array")}, + routes=[RouteDef(to="fans")], + ), + ], + for_each=[ + ForEachDef( + name="fans", + type="for_each", + source="lister.output.repos", + **{"as": "repo"}, + agent=AgentDef( + name="fan_agent", + model="gpt-4", + prompt="Work {{ repo }}", + output={"r": OutputField(type="string")}, + ), + routes=[RouteDef(to="$end")], + ), + ], + output={}, + ) + provider = _RecordingWorkingDirProvider() + + async def _execute(agent, context, rendered_prompt, tools=None, **kwargs): + if agent.name == "lister": + return AgentOutput( + content={"repos": ["one"]}, + raw_response=None, + model=agent.model, + input_tokens=1, + output_tokens=1, + ) + return AgentOutput( + content={"r": "ok"}, + raw_response=None, + model=agent.model, + input_tokens=1, + output_tokens=1, + ) + + provider.execute = _execute # type: ignore[method-assign] + events: list = [] + emitter = WorkflowEventEmitter() + emitter.subscribe(events.append) + engine = WorkflowEngine( + config, provider, event_emitter=emitter, workflow_path=_workflow_file(tmp_path) + ) + + await engine.run({}) + + envelope = [e for e in events if e.type == "for_each_item_started"] + assert len(envelope) == 1 + assert envelope[0].data == {"group_name": "fans", "item_key": "0", "index": 0} diff --git a/tests/test_integration/test_claude_mcp_tool_filter.py b/tests/test_integration/test_claude_mcp_tool_filter.py index a6398897..5fa74dac 100644 --- a/tests/test_integration/test_claude_mcp_tool_filter.py +++ b/tests/test_integration/test_claude_mcp_tool_filter.py @@ -14,6 +14,7 @@ from __future__ import annotations +import os from typing import Any from unittest.mock import AsyncMock, MagicMock @@ -86,14 +87,14 @@ def _make_provider_with_mcp() -> ClaudeProvider: provider._default_max_session_seconds = None provider._default_reasoning_effort = None - # Pre-wire a mock MCP manager so _ensure_mcp_connected is a no-op mock_mcp = MagicMock() mock_mcp.get_all_tools.return_value = FAKE_MCP_TOOLS mock_mcp.has_servers.return_value = True - provider._mcp_manager = mock_mcp provider._mcp_servers_config = { "filesystem": {"command": "npx", "args": ["-y", "@modelcontextprotocol/server-filesystem"]}, } + provider._mcp_managers = {os.getcwd(): mock_mcp} + provider._mcp_manager_locks = {} return provider diff --git a/tests/test_mcp/test_manager.py b/tests/test_mcp/test_manager.py index e02bc9a9..af2ce790 100644 --- a/tests/test_mcp/test_manager.py +++ b/tests/test_mcp/test_manager.py @@ -277,3 +277,97 @@ async def test_connect_server_mocked(self, manager: Any) -> None: assert "web-search" in manager.sessions assert "web-search" in manager.tools assert manager.tool_to_server["web-search__search"] == "web-search" + + async def test_connect_server_forwards_cwd(self, manager: Any) -> None: + """Requirement: agent working_dir must reach the MCP server subprocess. + + ``connect_server(cwd=...)`` must forward the value to + ``StdioServerParameters(cwd=...)`` so the stdio MCP server process + is spawned in the agent's working directory (issue: + agent-mcp-working-dir todo 4). + """ + mock_list_tools_response = MagicMock() + mock_list_tools_response.tools = [] + + mock_session = AsyncMock() + mock_session.initialize = AsyncMock() + mock_session.list_tools = AsyncMock(return_value=mock_list_tools_response) + + mock_read_stream = MagicMock() + mock_write_stream = MagicMock() + mock_stdio_context = MagicMock() + mock_client_session = MagicMock() + + with ( + patch("conductor.mcp.manager.MCP_SDK_AVAILABLE", True), + patch("conductor.mcp.manager.StdioServerParameters") as mock_params_cls, + patch( + "conductor.mcp.manager.stdio_client", + return_value=mock_stdio_context, + ), + patch( + "conductor.mcp.manager.ClientSession", + return_value=mock_client_session, + ), + ): + manager._exit_stack.enter_async_context = AsyncMock( + side_effect=[ + (mock_read_stream, mock_write_stream), + mock_session, + ] + ) + + await manager.connect_server( + name="fs", + command="npx", + args=["-y", "@modelcontextprotocol/server-filesystem"], + cwd="/repo/worktree-a", + ) + + # StdioServerParameters must receive the cwd so the child process + # starts in that directory. + mock_params_cls.assert_called_once() + assert mock_params_cls.call_args.kwargs["cwd"] == "/repo/worktree-a" + + async def test_connect_server_cwd_defaults_to_none(self, manager: Any) -> None: + """Requirement: omitting cwd keeps legacy behavior (spawn in process cwd). + + When no cwd is given, ``StdioServerParameters`` must still be built + with ``cwd=None`` so the MCP SDK spawns the server in the conductor + process's own working directory (backward compatible). + """ + mock_list_tools_response = MagicMock() + mock_list_tools_response.tools = [] + + mock_session = AsyncMock() + mock_session.initialize = AsyncMock() + mock_session.list_tools = AsyncMock(return_value=mock_list_tools_response) + + mock_read_stream = MagicMock() + mock_write_stream = MagicMock() + mock_stdio_context = MagicMock() + mock_client_session = MagicMock() + + with ( + patch("conductor.mcp.manager.MCP_SDK_AVAILABLE", True), + patch("conductor.mcp.manager.StdioServerParameters") as mock_params_cls, + patch( + "conductor.mcp.manager.stdio_client", + return_value=mock_stdio_context, + ), + patch( + "conductor.mcp.manager.ClientSession", + return_value=mock_client_session, + ), + ): + manager._exit_stack.enter_async_context = AsyncMock( + side_effect=[ + (mock_read_stream, mock_write_stream), + mock_session, + ] + ) + + await manager.connect_server(name="fs", command="npx") + + mock_params_cls.assert_called_once() + assert mock_params_cls.call_args.kwargs["cwd"] is None diff --git a/tests/test_providers/test_capabilities.py b/tests/test_providers/test_capabilities.py index 12bae939..5231cdc0 100644 --- a/tests/test_providers/test_capabilities.py +++ b/tests/test_providers/test_capabilities.py @@ -35,7 +35,7 @@ def _stable_capabilities(**overrides: object) -> ProviderCapabilities: class TestSchemaValidation: def test_construct_stable_descriptor(self) -> None: - caps = _stable_capabilities() + caps = _stable_capabilities(working_dir=True) assert caps.tier == "stable" assert caps.is_experimental is False assert caps.declared_limitations() == [] @@ -118,12 +118,23 @@ def test_invalid_structured_output_mode_rejected(self) -> None: with pytest.raises(ValueError): _stable_capabilities(structured_output="json_mode") + def test_working_dir_defaults_to_false(self) -> None: + """Requirement: ``working_dir`` capability defaults to False so existing + descriptors built without the field keep the conservative value.""" + caps = _stable_capabilities() + assert caps.working_dir is False + + def test_working_dir_field_accepts_bool(self) -> None: + """Requirement: ``working_dir`` is a declared capability field (frozen schema).""" + assert _stable_capabilities(working_dir=True).working_dir is True + assert _stable_capabilities(working_dir=False).working_dir is False + class TestDeclaredLimitations: """Auto-generated limitations line for the experimental banner.""" def test_fully_stable_has_no_limitations(self) -> None: - assert _stable_capabilities().declared_limitations() == [] + assert _stable_capabilities(working_dir=True).declared_limitations() == [] def test_each_false_flag_produces_a_limitation(self) -> None: caps = _stable_capabilities( @@ -189,6 +200,34 @@ def test_every_production_provider_has_capabilities(self, provider_name: str) -> caps = get_capabilities(provider_name) assert isinstance(caps, ProviderCapabilities) + @pytest.mark.parametrize( + ("provider_name", "expected"), + [ + # Requirement: copilot and claude honor agent/runtime ``working_dir`` + # for the SDK session and its MCP servers; hermes and + # claude-agent-sdk do not (declared False so validate errors out). + ("copilot", True), + ("claude", True), + ("hermes", False), + ("claude-agent-sdk", False), + ], + ) + def test_working_dir_capability_matrix(self, provider_name: str, expected: bool) -> None: + """Each production provider declares an accurate ``working_dir`` flag.""" + if provider_name == "claude-agent-sdk": + pytest.importorskip("claude_agent_sdk") + caps = get_capabilities(provider_name) + assert caps.working_dir is expected + + def test_working_dir_false_listed_as_limitation(self) -> None: + """Requirement: the experimental banner surfaces working_dir=False.""" + lims = _stable_capabilities(working_dir=False).declared_limitations() + assert "working_dir ignored" in lims + assert ( + "working_dir ignored" + not in _stable_capabilities(working_dir=True).declared_limitations() + ) + def test_resolver_does_not_instantiate_provider(self) -> None: """The validator runs without API keys, so resolution MUST be class-only. diff --git a/tests/test_providers/test_claude.py b/tests/test_providers/test_claude.py index fd319b54..672f5668 100644 --- a/tests/test_providers/test_claude.py +++ b/tests/test_providers/test_claude.py @@ -9,6 +9,8 @@ - Error handling and wrapping """ +import asyncio +import os from unittest.mock import AsyncMock, Mock, patch import pytest @@ -3124,10 +3126,12 @@ async def test_thinking_block_preserved_in_agentic_loop_replay( # Stub MCP manager so the tool_use is actually executed and the loop # advances to iteration 2 (without an MCP manager the loop bails out). + # Injected into the cwd pool under the current process cwd, which is + # what the agent (working_dir=None) resolves to. mock_mcp = Mock() mock_mcp.has_servers = Mock(return_value=False) # don't add tools to the request mock_mcp.call_tool = AsyncMock(return_value="sunny") - provider._mcp_manager = mock_mcp + provider._mcp_managers[os.getcwd()] = mock_mcp agent = AgentDef( name="t", @@ -3209,7 +3213,7 @@ async def test_redacted_thinking_block_preserved_in_agentic_loop_replay( mock_mcp = Mock() mock_mcp.has_servers = Mock(return_value=False) mock_mcp.call_tool = AsyncMock(return_value="ok") - provider._mcp_manager = mock_mcp + provider._mcp_managers[os.getcwd()] = mock_mcp agent = AgentDef( name="t", @@ -3403,3 +3407,299 @@ async def test_thinking_kwarg_forwarded_through_parse_recovery_path( # is a fourth messages.create site reachable only via mid-agent interrupt # and partial-output flow; mocking complexity is prohibitive for a unit # test (would need a full asyncio interrupt fixture). + + +class TestClaudeMCPManagerPool: + """Requirement: agents with different working_dirs must get isolated MCP servers. + + Covers the MCPManager pool keyed by resolved cwd (agent-mcp-working-dir + todo 4): two cwds → two managers, repeated cwd → reuse, close() closes + all, parallel agents with different cwds never share a manager, a + connect failure for one cwd does not break another (fail-open), and the + resolved cwd is forwarded to ``MCPManager.connect_server(cwd=...)``. + """ + + @staticmethod + def _build_provider( + mock_anthropic_module: Mock, + mock_anthropic_class: Mock, + mcp_servers: dict[str, dict[str, str]], + ) -> ClaudeProvider: + """Build a ClaudeProvider with a mocked Anthropic client and MCP config.""" + mock_anthropic_module.__version__ = "0.77.0" + mock_client = Mock() + mock_client.models.list = AsyncMock(return_value=Mock(data=[])) + mock_client.messages.create = AsyncMock(return_value=Mock(content=[])) + mock_client.close = AsyncMock() + mock_anthropic_class.return_value = mock_client + return ClaudeProvider(mcp_servers=mcp_servers) + + @staticmethod + def _manager_factory(instances: list[Mock]) -> type: + """Return a fake MCPManager class whose instances are tracked in ``instances``.""" + + class _FakeMCPManager: + def __init__(self) -> None: + self.connected: list[dict[str, object]] = [] + self.closed = False + instances.append(self) # type: ignore[arg-type] + + async def connect_server(self, **kwargs: object) -> list[dict[str, object]]: + # Yield so concurrent first-callers genuinely interleave inside + # the critical section — without this the gather runs them + # sequentially and the double-checked locking is never exercised. + await asyncio.sleep(0) + self.connected.append(kwargs) + return [] + + def has_servers(self) -> bool: + return len(self.connected) > 0 + + def get_all_tools(self) -> list[dict[str, object]]: + return [] + + async def call_tool(self, name: str, arguments: dict[str, object]) -> str: + return "ok" + + async def close(self) -> None: + self.closed = True + + return _FakeMCPManager # type: ignore[return-value] + + @patch("conductor.providers.claude.ANTHROPIC_SDK_AVAILABLE", True) + @patch("conductor.providers.claude.AsyncAnthropic") + @patch("conductor.providers.claude.anthropic") + @pytest.mark.asyncio + async def test_mcp_pool_two_cwds_create_two_managers( + self, mock_anthropic_module: Mock, mock_anthropic_class: Mock + ) -> None: + """Requirement: two distinct resolved cwds → two distinct pool entries. + + Each cwd needs its own MCPManager because stdio MCP servers are + spawned per-manager with that manager's cwd; sharing one manager + across cwds would silently run tools in the wrong directory. + """ + servers = {"fs": {"command": "npx", "args": []}} + provider = self._build_provider(mock_anthropic_module, mock_anthropic_class, servers) + + instances: list[Mock] = [] + fake_cls = self._manager_factory(instances) + with ( + patch("conductor.mcp.manager.MCP_SDK_AVAILABLE", True), + patch("conductor.mcp.manager.MCPManager", fake_cls), + ): + manager_a = await provider._get_mcp_manager_for_cwd("/repo/a") + manager_b = await provider._get_mcp_manager_for_cwd("/repo/b") + + assert manager_a is not manager_b + assert len(instances) == 2 + # The cwd must be forwarded to connect_server for each pool entry. + assert instances[0].connected[0]["cwd"] == "/repo/a" + assert instances[1].connected[0]["cwd"] == "/repo/b" + + @patch("conductor.providers.claude.ANTHROPIC_SDK_AVAILABLE", True) + @patch("conductor.providers.claude.AsyncAnthropic") + @patch("conductor.providers.claude.anthropic") + @pytest.mark.asyncio + async def test_mcp_pool_repeated_cwd_reuses_manager( + self, mock_anthropic_module: Mock, mock_anthropic_class: Mock + ) -> None: + """Requirement: repeated resolution of the same cwd reuses the manager. + + Spawning a fresh MCP server process per agent execution would be + prohibitively expensive; the pool must return the already-connected + manager for a cwd seen before (no per-agent spawn/teardown). + """ + servers = {"fs": {"command": "npx", "args": []}} + provider = self._build_provider(mock_anthropic_module, mock_anthropic_class, servers) + + instances: list[Mock] = [] + fake_cls = self._manager_factory(instances) + with ( + patch("conductor.mcp.manager.MCP_SDK_AVAILABLE", True), + patch("conductor.mcp.manager.MCPManager", fake_cls), + ): + first = await provider._get_mcp_manager_for_cwd("/repo/a") + second = await provider._get_mcp_manager_for_cwd("/repo/a") + + assert first is second + assert len(instances) == 1 + # connect_server ran exactly once per configured server (no reconnect). + assert len(instances[0].connected) == 1 + + @patch("conductor.providers.claude.ANTHROPIC_SDK_AVAILABLE", True) + @patch("conductor.providers.claude.AsyncAnthropic") + @patch("conductor.providers.claude.anthropic") + @pytest.mark.asyncio + async def test_mcp_pool_close_closes_all_managers( + self, mock_anthropic_module: Mock, mock_anthropic_class: Mock + ) -> None: + """Requirement: provider.close() closes every pooled manager (idempotent). + + MCP server subprocesses are tied to the provider lifetime; close() + must iterate the whole pool and shut down each manager, and a second + close() must be a safe no-op. + """ + servers = {"fs": {"command": "npx", "args": []}} + provider = self._build_provider(mock_anthropic_module, mock_anthropic_class, servers) + + instances: list[Mock] = [] + fake_cls = self._manager_factory(instances) + with ( + patch("conductor.mcp.manager.MCP_SDK_AVAILABLE", True), + patch("conductor.mcp.manager.MCPManager", fake_cls), + ): + await provider._get_mcp_manager_for_cwd("/repo/a") + await provider._get_mcp_manager_for_cwd("/repo/b") + await provider.close() + await provider.close() # idempotent: must not raise + + assert len(instances) == 2 + assert all(inst.closed for inst in instances) + + @patch("conductor.providers.claude.ANTHROPIC_SDK_AVAILABLE", True) + @patch("conductor.providers.claude.AsyncAnthropic") + @patch("conductor.providers.claude.anthropic") + @pytest.mark.asyncio + async def test_mcp_pool_parallel_agents_no_race( + self, mock_anthropic_module: Mock, mock_anthropic_class: Mock + ) -> None: + """Requirement: parallel agents with different cwds never share a manager. + + Concurrent first-use of the pool must not race: each cwd ends up with + exactly one manager, and two agents resolving the same cwd observe the + same instance (lock-guarded lazy connect). + """ + servers = {"fs": {"command": "npx", "args": []}} + provider = self._build_provider(mock_anthropic_module, mock_anthropic_class, servers) + + instances: list[Mock] = [] + fake_cls = self._manager_factory(instances) + with ( + patch("conductor.mcp.manager.MCP_SDK_AVAILABLE", True), + patch("conductor.mcp.manager.MCPManager", fake_cls), + ): + results = await asyncio.gather( + provider._get_mcp_manager_for_cwd("/repo/a"), + provider._get_mcp_manager_for_cwd("/repo/b"), + provider._get_mcp_manager_for_cwd("/repo/a"), + provider._get_mcp_manager_for_cwd("/repo/b"), + ) + + # Exactly two managers total — one per cwd — despite concurrent access. + assert len(instances) == 2 + assert results[0] is results[2] + assert results[1] is results[3] + assert results[0] is not results[1] + # The lock serialized the lazy connect: exactly one connect per cwd + # (the second caller for each cwd hit the re-check under the lock). + per_cwd_connects = sorted(kwargs["cwd"] for i in instances for kwargs in i.connected) + assert per_cwd_connects == ["/repo/a", "/repo/b"] + + @patch("conductor.providers.claude.ANTHROPIC_SDK_AVAILABLE", True) + @patch("conductor.providers.claude.AsyncAnthropic") + @patch("conductor.providers.claude.anthropic") + @pytest.mark.asyncio + async def test_mcp_pool_connect_failure_fail_open_per_cwd( + self, mock_anthropic_module: Mock, mock_anthropic_class: Mock + ) -> None: + """Requirement: a connect failure for one cwd must not break another cwd. + + Fail-open per-server connect is existing behavior (the error is + logged and remaining servers still connect). The pool must preserve + it across pool keys: cwd A failing to connect leaves cwd B fully + functional. A cwd where EVERY server failed is not pooled, so a + later call for that same cwd retries the connect and can succeed. + """ + servers = {"fs": {"command": "npx", "args": []}} + provider = self._build_provider(mock_anthropic_module, mock_anthropic_class, servers) + + instances: list[Mock] = [] + fake_cls = self._manager_factory(instances) + attempts: dict[str, int] = {} + + async def failing_connect(self: Mock, **kwargs: object) -> list[dict[str, object]]: + cwd = kwargs.get("cwd") + assert isinstance(cwd, str) + attempts[cwd] = attempts.get(cwd, 0) + 1 + # Fail only the FIRST attempt for /repo/bad; the retry succeeds. + if cwd == "/repo/bad" and attempts[cwd] == 1: + raise RuntimeError("spawn failed") + self.connected.append(kwargs) + return [] + + with ( + patch("conductor.mcp.manager.MCP_SDK_AVAILABLE", True), + patch("conductor.mcp.manager.MCPManager", fake_cls), + patch.object(fake_cls, "connect_server", failing_connect), + ): + manager_bad = await provider._get_mcp_manager_for_cwd("/repo/bad") + manager_good = await provider._get_mcp_manager_for_cwd("/repo/good") + manager_bad_retry = await provider._get_mcp_manager_for_cwd("/repo/bad") + + # The failed first attempt left cwd /repo/bad with zero connections, + # so it was NOT pooled; the good cwd connected and was pooled. + assert manager_bad is not manager_good + assert instances[0].connected == [] # /repo/bad, first attempt: connect raised + assert instances[1].connected[0]["cwd"] == "/repo/good" + # The retry for /repo/bad built a NEW manager (the first was not + # cached) and re-ran the connect, which succeeded this time. + assert manager_bad_retry is not manager_bad + assert len(instances) == 3 + assert instances[2].connected[0]["cwd"] == "/repo/bad" + assert attempts["/repo/bad"] == 2 + + @patch("conductor.providers.claude.ANTHROPIC_SDK_AVAILABLE", True) + @patch("conductor.providers.claude.AsyncAnthropic") + @patch("conductor.providers.claude.anthropic") + @pytest.mark.asyncio + async def test_mcp_pool_no_config_returns_none( + self, mock_anthropic_module: Mock, mock_anthropic_class: Mock + ) -> None: + """Requirement: no runtime.mcp_servers configured → no pool entries. + + When the workflow declares no MCP servers the helper must return None + and never construct an MCPManager, regardless of the agent's cwd. + """ + provider = self._build_provider(mock_anthropic_module, mock_anthropic_class, {}) + + with patch("conductor.mcp.manager.MCP_SDK_AVAILABLE", True): + result = await provider._get_mcp_manager_for_cwd("/repo/a") + + assert result is None + + @patch("conductor.providers.claude.ANTHROPIC_SDK_AVAILABLE", True) + @patch("conductor.providers.claude.AsyncAnthropic") + @patch("conductor.providers.claude.anthropic") + @pytest.mark.asyncio + async def test_mcp_pool_agent_without_working_dir_uses_process_cwd( + self, + mock_anthropic_module: Mock, + mock_anthropic_class: Mock, + tmp_path: object, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + """Requirement: agent.working_dir=None falls back to os.getcwd() as pool key. + + Agents that do not declare working_dir must behave exactly as before + the pool existed: their MCP servers spawn in the conductor process's + current working directory (precedence agent > runtime > cwd). + """ + workdir = str(tmp_path) + monkeypatch.chdir(workdir) + + servers = {"fs": {"command": "npx", "args": []}} + provider = self._build_provider(mock_anthropic_module, mock_anthropic_class, servers) + + instances: list[Mock] = [] + fake_cls = self._manager_factory(instances) + with ( + patch("conductor.mcp.manager.MCP_SDK_AVAILABLE", True), + patch("conductor.mcp.manager.MCPManager", fake_cls), + ): + agent = AgentDef(name="a", prompt="p") # working_dir=None + resolved_cwd = agent.working_dir or os.getcwd() + manager = await provider._get_mcp_manager_for_cwd(resolved_cwd) + + assert manager is not None + assert instances[0].connected[0]["cwd"] == os.getcwd() diff --git a/tests/test_providers/test_claude_event_callback.py b/tests/test_providers/test_claude_event_callback.py index 02dd2564..df2a0ef5 100644 --- a/tests/test_providers/test_claude_event_callback.py +++ b/tests/test_providers/test_claude_event_callback.py @@ -69,7 +69,7 @@ def _make_provider_with_mcp() -> ClaudeProvider: mock_mcp_manager.has_servers.return_value = True mock_mcp_manager.get_all_tools.return_value = [] mock_mcp_manager.call_tool = AsyncMock(return_value="tool result") - provider._mcp_manager = mock_mcp_manager + provider._mock_mcp_manager = mock_mcp_manager return provider @@ -78,7 +78,6 @@ def _make_bare_provider() -> ClaudeProvider: """Create a minimal ClaudeProvider without MCP.""" provider = ClaudeProvider.__new__(ClaudeProvider) provider._client = MagicMock() - provider._mcp_manager = None provider._mcp_servers_config = None provider._default_model = "claude-3-5-sonnet-latest" provider._default_temperature = None @@ -120,6 +119,7 @@ async def test_emitted_on_single_iteration(self) -> None: output_schema=None, has_output_schema=False, event_callback=lambda t, d: events.append((t, d)), + mcp_manager=getattr(provider, "_mock_mcp_manager", None), ) turn_events = [(t, d) for t, d in events if t == "agent_turn_start"] @@ -153,6 +153,7 @@ async def test_emitted_on_multiple_iterations(self) -> None: output_schema=None, has_output_schema=False, event_callback=lambda t, d: events.append((t, d)), + mcp_manager=getattr(provider, "_mock_mcp_manager", None), ) turn_events = [(t, d) for t, d in events if t == "agent_turn_start"] @@ -189,6 +190,7 @@ async def test_emitted_for_text_response(self) -> None: output_schema=None, has_output_schema=False, event_callback=lambda t, d: events.append((t, d)), + mcp_manager=getattr(provider, "_mock_mcp_manager", None), ) msg_events = [(t, d) for t, d in events if t == "agent_message"] @@ -218,6 +220,7 @@ async def test_multiple_text_blocks(self) -> None: output_schema=None, has_output_schema=False, event_callback=lambda t, d: events.append((t, d)), + mcp_manager=getattr(provider, "_mock_mcp_manager", None), ) msg_events = [(t, d) for t, d in events if t == "agent_message"] @@ -243,6 +246,7 @@ async def test_empty_text_not_emitted(self) -> None: output_schema=None, has_output_schema=False, event_callback=lambda t, d: events.append((t, d)), + mcp_manager=getattr(provider, "_mock_mcp_manager", None), ) msg_events = [(t, d) for t, d in events if t == "agent_message"] @@ -270,7 +274,7 @@ async def test_tool_start_and_complete_on_success(self) -> None: ) text_response = _make_response([_make_text_block("Done")]) provider._execute_api_call = AsyncMock(side_effect=[mcp_response, text_response]) - provider._mcp_manager.call_tool = AsyncMock(return_value="file content here") + provider._mock_mcp_manager.call_tool = AsyncMock(return_value="file content here") await provider._execute_agentic_loop( messages=[{"role": "user", "content": "test"}], @@ -281,6 +285,7 @@ async def test_tool_start_and_complete_on_success(self) -> None: output_schema=None, has_output_schema=False, event_callback=lambda t, d: events.append((t, d)), + mcp_manager=getattr(provider, "_mock_mcp_manager", None), ) start_events = [(t, d) for t, d in events if t == "agent_tool_start"] @@ -306,7 +311,7 @@ async def test_tool_complete_on_failure(self) -> None: ) text_response = _make_response([_make_text_block("Error handled")]) provider._execute_api_call = AsyncMock(side_effect=[mcp_response, text_response]) - provider._mcp_manager.call_tool = AsyncMock(side_effect=RuntimeError("File not found")) + provider._mock_mcp_manager.call_tool = AsyncMock(side_effect=RuntimeError("File not found")) await provider._execute_agentic_loop( messages=[{"role": "user", "content": "test"}], @@ -317,6 +322,7 @@ async def test_tool_complete_on_failure(self) -> None: output_schema=None, has_output_schema=False, event_callback=lambda t, d: events.append((t, d)), + mcp_manager=getattr(provider, "_mock_mcp_manager", None), ) start_events = [(t, d) for t, d in events if t == "agent_tool_start"] @@ -343,7 +349,7 @@ async def test_multiple_tool_calls(self) -> None: ) text_response = _make_response([_make_text_block("Done")]) provider._execute_api_call = AsyncMock(side_effect=[mcp_response, text_response]) - provider._mcp_manager.call_tool = AsyncMock(return_value="result") + provider._mock_mcp_manager.call_tool = AsyncMock(return_value="result") await provider._execute_agentic_loop( messages=[{"role": "user", "content": "test"}], @@ -354,6 +360,7 @@ async def test_multiple_tool_calls(self) -> None: output_schema=None, has_output_schema=False, event_callback=lambda t, d: events.append((t, d)), + mcp_manager=getattr(provider, "_mock_mcp_manager", None), ) start_events = [(t, d) for t, d in events if t == "agent_tool_start"] @@ -388,6 +395,7 @@ async def test_tool_arguments_truncated(self) -> None: output_schema=None, has_output_schema=False, event_callback=lambda t, d: events.append((t, d)), + mcp_manager=getattr(provider, "_mock_mcp_manager", None), ) start_events = [(t, d) for t, d in events if t == "agent_tool_start"] @@ -409,7 +417,7 @@ async def test_tool_result_truncated(self) -> None: ) text_response = _make_response([_make_text_block("Done")]) provider._execute_api_call = AsyncMock(side_effect=[mcp_response, text_response]) - provider._mcp_manager.call_tool = AsyncMock(return_value="y" * 600) + provider._mock_mcp_manager.call_tool = AsyncMock(return_value="y" * 600) await provider._execute_agentic_loop( messages=[{"role": "user", "content": "test"}], @@ -420,6 +428,7 @@ async def test_tool_result_truncated(self) -> None: output_schema=None, has_output_schema=False, event_callback=lambda t, d: events.append((t, d)), + mcp_manager=getattr(provider, "_mock_mcp_manager", None), ) complete_events = [(t, d) for t, d in events if t == "agent_tool_complete"] @@ -454,6 +463,7 @@ async def test_text_response_no_callback(self) -> None: output_schema=None, has_output_schema=False, event_callback=None, + mcp_manager=getattr(provider, "_mock_mcp_manager", None), ) assert response is text_response @@ -481,6 +491,7 @@ async def test_tool_calls_no_callback(self) -> None: output_schema=None, has_output_schema=False, event_callback=None, + mcp_manager=getattr(provider, "_mock_mcp_manager", None), ) assert response is text_response @@ -520,6 +531,7 @@ def bad_callback(event_type: str, data: dict[str, Any]) -> None: output_schema=None, has_output_schema=False, event_callback=bad_callback, + mcp_manager=getattr(provider, "_mock_mcp_manager", None), ) # Should still return normally despite callback errors @@ -544,8 +556,8 @@ async def test_callback_reaches_agentic_loop(self) -> None: text_response = _make_response([_make_text_block("Hello")]) provider._execute_api_call = AsyncMock(return_value=text_response) - # Mock _ensure_mcp_connected since we're calling _execute_with_retry directly - provider._ensure_mcp_connected = AsyncMock() + # Mock the cwd-pool lookup since we're calling _execute_with_retry directly + provider._get_mcp_manager_for_cwd = AsyncMock(return_value=None) # Create a minimal agent mock agent = MagicMock() @@ -607,7 +619,7 @@ async def test_full_sequence(self) -> None: ) provider._execute_api_call = AsyncMock(side_effect=[turn1_response, turn2_response]) - provider._mcp_manager.call_tool = AsyncMock(return_value="hello") + provider._mock_mcp_manager.call_tool = AsyncMock(return_value="hello") await provider._execute_agentic_loop( messages=[{"role": "user", "content": "read /tmp/test.txt"}], @@ -618,6 +630,7 @@ async def test_full_sequence(self) -> None: output_schema=None, has_output_schema=False, event_callback=lambda t, d: events.append((t, d)), + mcp_manager=getattr(provider, "_mock_mcp_manager", None), ) event_types = [t for t, _ in events] diff --git a/tests/test_providers/test_claude_interrupt.py b/tests/test_providers/test_claude_interrupt.py index 9c9f7a75..4c9cb432 100644 --- a/tests/test_providers/test_claude_interrupt.py +++ b/tests/test_providers/test_claude_interrupt.py @@ -26,7 +26,7 @@ def _make_provider() -> ClaudeProvider: """Create a ClaudeProvider with essential attributes for testing.""" provider = ClaudeProvider.__new__(ClaudeProvider) provider._client = MagicMock() - provider._mcp_manager = None + provider._mcp_managers = {} provider._mcp_servers_config = None provider._default_model = "claude-3-5-sonnet-latest" provider._default_temperature = None @@ -187,8 +187,8 @@ async def capture_api_call(messages: Any, **kwargs: Any) -> MagicMock: async def test_interrupt_on_second_iteration(self) -> None: """Interrupt is detected after the first iteration's tool call completes.""" provider = _make_provider() - provider._mcp_manager = MagicMock() - provider._mcp_manager.call_tool = AsyncMock(return_value="tool result") + mock_mcp_manager = MagicMock() + mock_mcp_manager.call_tool = AsyncMock(return_value="tool result") interrupt = asyncio.Event() @@ -224,6 +224,7 @@ async def mock_parse_recovery(messages: Any, **kwargs: Any) -> MagicMock: output_schema={"result": OutputField(type="string")}, has_output_schema=True, interrupt_signal=interrupt, + mcp_manager=mock_mcp_manager, ) assert is_partial is True diff --git a/tests/test_providers/test_claude_mcp_tool_filter.py b/tests/test_providers/test_claude_mcp_tool_filter.py index e8a45cce..513b9ad0 100644 --- a/tests/test_providers/test_claude_mcp_tool_filter.py +++ b/tests/test_providers/test_claude_mcp_tool_filter.py @@ -38,7 +38,7 @@ def _make_provider_with_mcp_tools(tools: list[dict[str, Any]]) -> ClaudeProvider mock_mcp_manager = MagicMock() mock_mcp_manager.get_all_tools.return_value = tools mock_mcp_manager.has_servers.return_value = True - provider._mcp_manager = mock_mcp_manager + provider._mock_mcp_manager = mock_mcp_manager return provider @@ -68,7 +68,9 @@ class TestConvertMcpToolsFilter: def test_none_filter_includes_all_tools(self) -> None: """tool_filter=None should include all MCP tools (no filtering).""" provider = _make_provider_with_mcp_tools(FAKE_MCP_TOOLS) - result = provider._convert_mcp_tools_to_claude(tool_filter=None) + result = provider._convert_mcp_tools_to_claude( + tool_filter=None, manager=provider._mock_mcp_manager + ) assert len(result) == 3 def test_empty_list_filter_includes_all_tools(self) -> None: @@ -78,13 +80,17 @@ def test_empty_list_filter_includes_all_tools(self) -> None: 'filter to nothing' instead of 'no filter applied'. """ provider = _make_provider_with_mcp_tools(FAKE_MCP_TOOLS) - result = provider._convert_mcp_tools_to_claude(tool_filter=[]) + result = provider._convert_mcp_tools_to_claude( + tool_filter=[], manager=provider._mock_mcp_manager + ) assert len(result) == 3 def test_specific_filter_includes_only_matching_tools(self) -> None: """tool_filter=['name'] should include only matching tools.""" provider = _make_provider_with_mcp_tools(FAKE_MCP_TOOLS) - result = provider._convert_mcp_tools_to_claude(tool_filter=["filesystem__read_file"]) + result = provider._convert_mcp_tools_to_claude( + tool_filter=["filesystem__read_file"], manager=provider._mock_mcp_manager + ) assert len(result) == 1 assert result[0]["name"] == "filesystem__read_file" @@ -92,7 +98,8 @@ def test_filter_with_multiple_tools(self) -> None: """tool_filter with multiple names should include all matching tools.""" provider = _make_provider_with_mcp_tools(FAKE_MCP_TOOLS) result = provider._convert_mcp_tools_to_claude( - tool_filter=["filesystem__read_file", "web_search__search"] + tool_filter=["filesystem__read_file", "web_search__search"], + manager=provider._mock_mcp_manager, ) assert len(result) == 2 names = {t["name"] for t in result} @@ -101,22 +108,25 @@ def test_filter_with_multiple_tools(self) -> None: def test_filter_with_nonexistent_tool_excludes_it(self) -> None: """tool_filter with names not in MCP tools should return no matches for those.""" provider = _make_provider_with_mcp_tools(FAKE_MCP_TOOLS) - result = provider._convert_mcp_tools_to_claude(tool_filter=["nonexistent_tool"]) + result = provider._convert_mcp_tools_to_claude( + tool_filter=["nonexistent_tool"], manager=provider._mock_mcp_manager + ) assert len(result) == 0 def test_no_mcp_manager_returns_empty(self) -> None: """No MCP manager should return empty list regardless of filter.""" provider = ClaudeProvider.__new__(ClaudeProvider) - provider._mcp_manager = None - result = provider._convert_mcp_tools_to_claude(tool_filter=[]) + result = provider._convert_mcp_tools_to_claude(tool_filter=[], manager=None) assert result == [] - result = provider._convert_mcp_tools_to_claude(tool_filter=None) + result = provider._convert_mcp_tools_to_claude(tool_filter=None, manager=None) assert result == [] def test_tool_format_preserved(self) -> None: """Converted tools should have name, description, and input_schema.""" provider = _make_provider_with_mcp_tools(FAKE_MCP_TOOLS) - result = provider._convert_mcp_tools_to_claude(tool_filter=None) + result = provider._convert_mcp_tools_to_claude( + tool_filter=None, manager=provider._mock_mcp_manager + ) for tool in result: assert "name" in tool assert "description" in tool @@ -162,7 +172,6 @@ def _make_bare_provider() -> ClaudeProvider: """Create a minimal ClaudeProvider for unit-testing helper methods.""" provider = ClaudeProvider.__new__(ClaudeProvider) provider._client = MagicMock() - provider._mcp_manager = None provider._mcp_servers_config = None provider._default_model = "claude-3-5-sonnet-latest" provider._default_temperature = None diff --git a/tests/test_providers/test_claude_system_prompt.py b/tests/test_providers/test_claude_system_prompt.py index 8f0df0ca..3858ff46 100644 --- a/tests/test_providers/test_claude_system_prompt.py +++ b/tests/test_providers/test_claude_system_prompt.py @@ -4,6 +4,7 @@ import asyncio import contextlib +import os from types import SimpleNamespace from typing import TYPE_CHECKING from unittest.mock import AsyncMock, Mock, patch @@ -104,7 +105,7 @@ async def test_tool_use_loop_passes_system_prompt_on_every_sdk_call( } ] mock_mcp.call_tool = AsyncMock(return_value="tool result") - provider._mcp_manager = mock_mcp + provider._mcp_managers[os.getcwd()] = mock_mcp agent = AgentDef.model_validate( { diff --git a/tests/test_providers/test_copilot_resume.py b/tests/test_providers/test_copilot_resume.py index 56349d72..0818af2c 100644 --- a/tests/test_providers/test_copilot_resume.py +++ b/tests/test_providers/test_copilot_resume.py @@ -9,6 +9,7 @@ from __future__ import annotations import logging +import os from typing import Any from unittest.mock import AsyncMock, MagicMock, patch @@ -164,9 +165,72 @@ async def _fake_send(msg: Any) -> None: mock_client.resume_session.assert_called_once_with( resumed_sid, on_permission_request=CopilotProvider._default_permission_handler, + working_directory=os.getcwd(), ) mock_client.create_session.assert_not_called() + @pytest.mark.asyncio + async def test_legacy_checkpoint_resume_logs_info( + self, caplog: pytest.LogCaptureFixture + ) -> None: + """Legacy checkpoint (no recorded cwd) resumes by ID and logs an INFO. + + Requirement: a checkpoint saved before working_dir tracking was + introduced has no recorded cwd. Resuming it must preserve the legacy + resume-by-id behavior (resume_session called with the resolved cwd, + no create_session fallback) and emit an INFO log explaining that the + checkpoint predates working_dir tracking. + """ + resumed_sid = "sess-old" + + mock_session = AsyncMock() + mock_session.session_id = "sess-resumed" + mock_session.disconnect = AsyncMock() + + def _fake_on(callback: Any) -> None: + mock_session._callback = callback + + mock_session.on = _fake_on + + async def _fake_send(msg: Any) -> None: + evt = MagicMock() + evt.type = MagicMock() + evt.type.value = "session.idle" + mock_session._callback(evt) + + mock_session.send = _fake_send + + mock_client = AsyncMock() + mock_client.resume_session = AsyncMock(return_value=mock_session) + mock_client.create_session = AsyncMock() # Should NOT be called + + provider = CopilotProvider() + provider._client = mock_client + provider._started = True + # Resume ID set, but NO recorded cwd — simulates a pre-cwd checkpoint. + provider.set_resume_session_ids({"researcher": resumed_sid}) + + agent = _make_agent("researcher") + + with ( + caplog.at_level(logging.INFO, logger="conductor.providers.copilot"), + patch("conductor.cli.app.is_verbose", return_value=False), + patch("conductor.cli.app.is_full", return_value=False), + ): + await provider.execute(agent, {}, "Continue research") + + mock_client.resume_session.assert_called_once_with( + resumed_sid, + on_permission_request=CopilotProvider._default_permission_handler, + working_directory=os.getcwd(), + ) + mock_client.create_session.assert_not_called() + assert any( + "without a recorded working directory" in r.message + and "predates working_dir tracking" in r.message + for r in caplog.records + ) + @pytest.mark.asyncio async def test_fallback_to_create_on_resume_runtime_error(self) -> None: """When resume_session raises RuntimeError, falls back to create_session.""" @@ -207,6 +271,7 @@ async def _fake_send(msg: Any) -> None: mock_client.resume_session.assert_called_once_with( "stale-sid", on_permission_request=CopilotProvider._default_permission_handler, + working_directory=os.getcwd(), ) mock_client.create_session.assert_called_once() # Session ID should now reflect the new session diff --git a/tests/test_providers/test_copilot_working_dir.py b/tests/test_providers/test_copilot_working_dir.py new file mode 100644 index 00000000..7610772a --- /dev/null +++ b/tests/test_providers/test_copilot_working_dir.py @@ -0,0 +1,322 @@ +"""Unit tests for per-agent working_directory support in the Copilot provider. + +Tests cover: +- ``session_kwargs["working_directory"]`` reflects the engine-resolved + ``agent.working_dir`` (falling back to ``os.getcwd()`` when unset). +- Stdio/local MCP server configs are stamped with ``working_directory`` per + execution without mutating the shared ``self._mcp_servers`` dict. +- HTTP/SSE MCP server configs are never stamped. +- Session resume forwards ``working_directory`` + stamped ``mcp_servers``. +- A changed working directory skips resume and creates a fresh session. +- ``get_session_cwds`` / ``set_resume_session_cwds`` tracking used by + checkpoint persistence. +""" + +from __future__ import annotations + +import copy +import logging +import os +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from conductor.config.schema import AgentDef +from conductor.providers.copilot import CopilotProvider + + +def _make_agent(name: str = "test_agent", working_dir: str | None = None) -> AgentDef: + """Create a minimal AgentDef for testing.""" + return AgentDef(name=name, model="gpt-4o", prompt="Test prompt", working_dir=working_dir) + + +def _make_idle_session(session_id: str) -> AsyncMock: + """Build a fake SDK session that resolves _send_and_wait immediately.""" + mock_session = AsyncMock() + mock_session.session_id = session_id + mock_session.disconnect = AsyncMock() + + def _fake_on(callback: Any) -> None: + mock_session._callback = callback + + mock_session.on = _fake_on + + async def _fake_send(msg: Any) -> None: + evt = MagicMock() + evt.type = MagicMock() + evt.type.value = "session.idle" + mock_session._callback(evt) + + mock_session.send = _fake_send + return mock_session + + +def _make_provider_with_client( + create_session: AsyncMock | None = None, + resume_session: AsyncMock | None = None, + mcp_servers: dict[str, Any] | None = None, +) -> CopilotProvider: + """Build a provider wired to a fake SDK client (SDK path, no mock_handler).""" + provider = CopilotProvider(mcp_servers=mcp_servers) + mock_client = AsyncMock() + mock_client.create_session = create_session or AsyncMock() + mock_client.resume_session = resume_session or AsyncMock() + provider._client = mock_client + provider._started = True + return provider + + +async def _execute(provider: CopilotProvider, agent: AgentDef) -> None: + """Run provider.execute with verbose helpers patched out.""" + with ( + patch("conductor.cli.app.is_verbose", return_value=False), + patch("conductor.cli.app.is_full", return_value=False), + ): + await provider.execute(agent, {}, "Do work") + + +def _capture_create() -> tuple[AsyncMock, dict[str, Any]]: + """Return (create_session mock, captured-kwargs dict).""" + captured: dict[str, Any] = {} + + async def _create(**kwargs: Any) -> AsyncMock: + captured.update(kwargs) + return _make_idle_session("sess-new") + + return AsyncMock(side_effect=_create), captured + + +# --------------------------------------------------------------------------- +# Session working_directory +# --------------------------------------------------------------------------- + + +class TestSessionWorkingDirectory: + """session_kwargs["working_directory"] must follow the resolved agent cwd.""" + + @pytest.mark.asyncio + async def test_working_directory_uses_agent_working_dir(self, tmp_path: Any) -> None: + """Requirement: session working_directory equals resolved agent.working_dir.""" + create_session, captured = _capture_create() + provider = _make_provider_with_client(create_session=create_session) + + await _execute(provider, _make_agent("a1", working_dir=str(tmp_path))) + + create_session.assert_called_once() + assert captured["working_directory"] == str(tmp_path) + + @pytest.mark.asyncio + async def test_working_directory_falls_back_to_os_getcwd(self) -> None: + """Requirement: working_directory is os.getcwd() when agent.working_dir is None.""" + create_session, captured = _capture_create() + provider = _make_provider_with_client(create_session=create_session) + + await _execute(provider, _make_agent("a1", working_dir=None)) + + assert captured["working_directory"] == os.getcwd() + + @pytest.mark.asyncio + async def test_session_cwd_tracked_for_checkpoint(self, tmp_path: Any) -> None: + """Requirement: provider tracks resolved cwd per agent for checkpoint persistence.""" + create_session, _ = _capture_create() + provider = _make_provider_with_client(create_session=create_session) + + await _execute(provider, _make_agent("a1", working_dir=str(tmp_path))) + + assert provider.get_session_cwds() == {"a1": str(tmp_path)} + + +# --------------------------------------------------------------------------- +# MCP server stamping +# --------------------------------------------------------------------------- + + +class TestMcpServerWorkingDirectoryStamping: + """Per-execution copies of stdio MCP configs get working_directory stamped.""" + + @pytest.mark.asyncio + async def test_stdio_server_stamped_with_working_directory(self, tmp_path: Any) -> None: + """Requirement: stdio MCP config carries working_directory == resolved cwd.""" + mcp_servers = { + "fs": {"type": "stdio", "command": "npx", "args": ["server"], "tools": "*"}, + } + create_session, captured = _capture_create() + provider = _make_provider_with_client( + create_session=create_session, mcp_servers=mcp_servers + ) + + await _execute(provider, _make_agent("a1", working_dir=str(tmp_path))) + + stamped = captured["mcp_servers"]["fs"] + assert stamped["working_directory"] == str(tmp_path) + + @pytest.mark.asyncio + async def test_http_and_sse_servers_not_stamped(self, tmp_path: Any) -> None: + """Requirement: http/sse MCP configs must NOT receive working_directory.""" + mcp_servers = { + "remote_http": {"type": "http", "url": "https://example.com/mcp", "tools": "*"}, + "remote_sse": {"type": "sse", "url": "https://example.com/sse", "tools": "*"}, + "local": {"type": "stdio", "command": "server", "tools": "*"}, + } + create_session, captured = _capture_create() + provider = _make_provider_with_client( + create_session=create_session, mcp_servers=mcp_servers + ) + + await _execute(provider, _make_agent("a1", working_dir=str(tmp_path))) + + assert "working_directory" not in captured["mcp_servers"]["remote_http"] + assert "working_directory" not in captured["mcp_servers"]["remote_sse"] + assert captured["mcp_servers"]["local"]["working_directory"] == str(tmp_path) + + @pytest.mark.asyncio + async def test_local_type_server_stamped(self, tmp_path: Any) -> None: + """Requirement: type "local" (SDK alias for stdio) is stamped as well.""" + mcp_servers = { + "fs": {"type": "local", "command": "server", "tools": "*"}, + } + create_session, captured = _capture_create() + provider = _make_provider_with_client( + create_session=create_session, mcp_servers=mcp_servers + ) + + await _execute(provider, _make_agent("a1", working_dir=str(tmp_path))) + + assert captured["mcp_servers"]["fs"]["working_directory"] == str(tmp_path) + + @pytest.mark.asyncio + async def test_shared_mcp_servers_dict_not_mutated(self, tmp_path: Any) -> None: + """Requirement: self._mcp_servers must remain unchanged after execute (deep-equal).""" + mcp_servers = { + "fs": {"type": "stdio", "command": "npx", "args": ["server"], "tools": "*"}, + "remote": {"type": "http", "url": "https://example.com/mcp", "tools": "*"}, + } + snapshot = copy.deepcopy(mcp_servers) + create_session, _ = _capture_create() + provider = _make_provider_with_client( + create_session=create_session, mcp_servers=mcp_servers + ) + + await _execute(provider, _make_agent("a1", working_dir=str(tmp_path))) + await _execute(provider, _make_agent("a2", working_dir=str(tmp_path))) + + assert provider._mcp_servers == snapshot + assert mcp_servers == snapshot + + +# --------------------------------------------------------------------------- +# Resume path +# --------------------------------------------------------------------------- + + +class TestResumeWorkingDirectory: + """resume_session receives working_directory + stamped mcp_servers.""" + + @pytest.mark.asyncio + async def test_resume_passes_working_directory_and_stamped_mcp(self, tmp_path: Any) -> None: + """Requirement: resume_session is called with resolved cwd and stamped stdio configs.""" + mcp_servers = { + "fs": {"type": "stdio", "command": "server", "tools": "*"}, + "remote": {"type": "http", "url": "https://example.com/mcp", "tools": "*"}, + } + resumed = _make_idle_session("sess-resumed") + resume_session = AsyncMock(return_value=resumed) + create_session = AsyncMock() + provider = _make_provider_with_client( + create_session=create_session, + resume_session=resume_session, + mcp_servers=mcp_servers, + ) + provider.set_resume_session_ids({"a1": "sid-old"}) + provider.set_resume_session_cwds({"a1": str(tmp_path)}) + + await _execute(provider, _make_agent("a1", working_dir=str(tmp_path))) + + resume_session.assert_called_once() + call = resume_session.call_args + assert call.args[0] == "sid-old" + assert call.kwargs["working_directory"] == str(tmp_path) + assert call.kwargs["mcp_servers"]["fs"]["working_directory"] == str(tmp_path) + assert "working_directory" not in call.kwargs["mcp_servers"]["remote"] + create_session.assert_not_called() + + @pytest.mark.asyncio + async def test_changed_cwd_skips_resume_and_creates_new_session( + self, tmp_path: Any, caplog: pytest.LogCaptureFixture + ) -> None: + """Requirement: resolved cwd != session-creation cwd => no resume, new session + warning.""" + other_dir = tmp_path / "other" + other_dir.mkdir() + resume_session = AsyncMock() + create_session, _ = _capture_create() + provider = _make_provider_with_client( + create_session=create_session, + resume_session=resume_session, + ) + provider.set_resume_session_ids({"a1": "sid-old"}) + provider.set_resume_session_cwds({"a1": str(tmp_path)}) + + with caplog.at_level(logging.WARNING, logger="conductor.providers.copilot"): + await _execute(provider, _make_agent("a1", working_dir=str(other_dir))) + + resume_session.assert_not_called() + create_session.assert_called_once() + assert any("working directory changed" in r.message for r in caplog.records) + + @pytest.mark.asyncio + async def test_unknown_previous_cwd_resumes_by_id(self, tmp_path: Any) -> None: + """Requirement: legacy checkpoints have no cwd record — resume by session id as before.""" + resumed = _make_idle_session("sess-resumed") + resume_session = AsyncMock(return_value=resumed) + create_session = AsyncMock() + provider = _make_provider_with_client( + create_session=create_session, + resume_session=resume_session, + ) + provider.set_resume_session_ids({"a1": "sid-old"}) + # No set_resume_session_cwds call — simulates a pre-cwd checkpoint. + + await _execute(provider, _make_agent("a1", working_dir=str(tmp_path))) + + resume_session.assert_called_once() + create_session.assert_not_called() + + @pytest.mark.asyncio + async def test_default_cwd_matches_tracked_cwd_resumes(self) -> None: + """Requirement: when resolved cwd equals the tracked cwd, resume proceeds.""" + resumed = _make_idle_session("sess-resumed") + resume_session = AsyncMock(return_value=resumed) + provider = _make_provider_with_client(resume_session=resume_session) + provider.set_resume_session_ids({"a1": "sid-old"}) + provider.set_resume_session_cwds({"a1": os.getcwd()}) + + await _execute(provider, _make_agent("a1", working_dir=None)) + + resume_session.assert_called_once() + assert resume_session.call_args.kwargs["working_directory"] == os.getcwd() + + +# --------------------------------------------------------------------------- +# Session cwd tracking API +# --------------------------------------------------------------------------- + + +class TestSessionCwdTracking: + """get/set resume cwd mapping used by checkpoint persistence.""" + + def test_set_resume_session_cwds_stores_copy(self) -> None: + """Requirement: set_resume_session_cwds stores a copy of the mapping.""" + provider = CopilotProvider(mock_handler=lambda a, p, c: {"r": 1}) + cwds = {"a1": "/repo/a"} + provider.set_resume_session_cwds(cwds) + cwds["a2"] = "/repo/b" + assert provider.get_session_cwds() == {} + + def test_get_session_cwds_returns_copy(self) -> None: + """Requirement: get_session_cwds returns a copy; mutations don't leak.""" + provider = CopilotProvider(mock_handler=lambda a, p, c: {"r": 1}) + provider._session_cwds["a1"] = "/repo/a" + result = provider.get_session_cwds() + result["a2"] = "/repo/b" + assert "a2" not in provider._session_cwds