diff --git a/src/cmdforge/providers.py b/src/cmdforge/providers.py index bea2483..8f01533 100644 --- a/src/cmdforge/providers.py +++ b/src/cmdforge/providers.py @@ -212,6 +212,8 @@ class Provider: fallback_chain: Optional[List[str]] = None # Ordered multi-step fallback (new) api_key_env: Optional[str] = None # Env var name for api-type providers pty_config: Optional[dict] = None # Patterns for pty-type providers + tools: Optional[List[str]] = None # Allowed CmdForge tools (None = all) + mcp_servers: Optional[List[str]] = None # MCP servers available to this provider def to_dict(self) -> dict: d = { @@ -236,6 +238,10 @@ class Provider: d["api_key_env"] = self.api_key_env if self.pty_config: d["pty_config"] = self.pty_config + if self.tools: + d["tools"] = self.tools + if self.mcp_servers: + d["mcp_servers"] = self.mcp_servers return d @classmethod @@ -252,6 +258,8 @@ class Provider: fallback_chain=data.get("fallback_chain"), api_key_env=data.get("api_key_env"), pty_config=data.get("pty_config"), + tools=data.get("tools"), + mcp_servers=data.get("mcp_servers"), ) diff --git a/src/cmdforge/runner.py b/src/cmdforge/runner.py index ca209fe..b2e37a9 100644 --- a/src/cmdforge/runner.py +++ b/src/cmdforge/runner.py @@ -5,7 +5,7 @@ import json import re import sys from pathlib import Path -from typing import Optional +from typing import Optional, List import yaml @@ -781,6 +781,23 @@ def execute_tool_step( # Determine effective provider (step override > parent override) effective_provider = step.provider or provider_override + # Check tool access if the provider restricts which tools it can call + if effective_provider: + from .providers import get_provider + provider_obj = get_provider(effective_provider) + if provider_obj and provider_obj.tools is not None: + allowed = provider_obj.tools + if step.tool not in allowed and not _tool_match_allowlist(step.tool, allowed): + _print_call_stack( + call_stack, + f"Provider '{effective_provider}' is not allowed to call tool '{step.tool}'", + ) + print( + f"Provider '{effective_provider}' tools allowlist: {allowed}", + file=sys.stderr, + ) + return "", False + if verbose: print(f"[verbose] Tool step: calling {step.tool}", file=sys.stderr) print(f"[verbose] Input length: {len(input_text)} chars", file=sys.stderr) @@ -799,7 +816,9 @@ def execute_tool_step( show_prompt=False, verbose=verbose, _depth=depth + 1, - _call_stack=call_stack + _call_stack=call_stack, + agent_profile=step.profile, + agent_skills=step.skills, ) return output, exit_code == 0 @@ -815,7 +834,9 @@ def run_tool( verbose: bool = False, auto_install: bool = False, _depth: int = 0, - _call_stack: Optional[list] = None + _call_stack: Optional[list] = None, + agent_profile: Optional[str] = None, + agent_skills: Optional[List[str]] = None, ) -> tuple[str, int]: """ Execute a tool. @@ -940,8 +961,14 @@ def run_tool( if dry_run: variables[step.output_var] = f"[DRY RUN - would call {step.provider}]" else: + # Apply agent-level defaults if step doesn't set its own + effective_step = step + if agent_profile and not step.profile: + effective_step = _step_with_profile(step, agent_profile) + if agent_skills is not None and step.skills is None: + effective_step = _step_with_skills(effective_step, agent_skills) output, success = execute_prompt_step( - step, + effective_step, variables, provider_override, verbose=verbose, @@ -1200,5 +1227,28 @@ def main(): sys.exit(exit_code) +def _tool_match_allowlist(name: str, allowed: list) -> bool: + for pattern in allowed: + if pattern == name: + return True + if pattern.endswith("*") and name.startswith(pattern[:-1]): + return True + return False + + +def _step_with_profile(step, profile_name): + from copy import copy + s = copy(step) + s.profile = profile_name + return s + + +def _step_with_skills(step, skill_names): + from copy import copy + s = copy(step) + s.skills = skill_names + return s + + if __name__ == "__main__": main() diff --git a/src/cmdforge/tool.py b/src/cmdforge/tool.py index 459fba8..5400d7b 100644 --- a/src/cmdforge/tool.py +++ b/src/cmdforge/tool.py @@ -241,6 +241,8 @@ class ToolStep: input_template: str = "{input}" # Input template (supports variable substitution) args: dict = field(default_factory=dict) # Arguments to pass to the tool provider: Optional[str] = None # Provider override for the called tool + profile: Optional[str] = None # AI persona for the nested tool + skills: Optional[List[str]] = None # Skills to enable in the nested tool name: Optional[str] = None # Optional display name for the step def to_dict(self) -> dict: @@ -255,6 +257,10 @@ class ToolStep: d["args"] = self.args if self.provider: d["provider"] = self.provider + if self.profile: + d["profile"] = self.profile + if self.skills is not None: + d["skills"] = self.skills if self.name: d["name"] = self.name return d @@ -267,6 +273,8 @@ class ToolStep: input_template=data.get("input", "{input}"), args=data.get("args", {}), provider=data.get("provider"), + profile=data.get("profile"), + skills=data.get("skills"), name=data.get("name") ) diff --git a/tests/test_providers.py b/tests/test_providers.py index 2df1ac3..7be65c7 100644 --- a/tests/test_providers.py +++ b/tests/test_providers.py @@ -784,3 +784,29 @@ class TestProviderDiscovery: assert len({item["name"] for item in models}) == 2 assert models[0]["name"] == "ollama-hermes4-3-latest" assert models[1]["name"] == "ollama-hermes4-3-q4-k-m" + + +class TestProviderAccessControl: + def test_tools_allowlist_defaults_to_none(self): + provider = Provider("test", "cmd") + assert provider.tools is None + assert provider.mcp_servers is None + + def test_tools_allowlist_roundtrip(self): + provider = Provider("test", "cmd", tools=["tool-a", "tool-b"]) + d = provider.to_dict() + assert d["tools"] == ["tool-a", "tool-b"] + restored = Provider.from_dict(d) + assert restored.tools == ["tool-a", "tool-b"] + + def test_mcp_servers_roundtrip(self): + provider = Provider("test", "cmd", mcp_servers=["filesystem"]) + d = provider.to_dict() + assert d["mcp_servers"] == ["filesystem"] + restored = Provider.from_dict(d) + assert restored.mcp_servers == ["filesystem"] + + def test_tools_not_in_dict_when_none(self): + provider = Provider("test", "cmd") + assert "tools" not in provider.to_dict() + assert "mcp_servers" not in provider.to_dict() diff --git a/tests/test_tool.py b/tests/test_tool.py index 920e5de..6856eb7 100644 --- a/tests/test_tool.py +++ b/tests/test_tool.py @@ -8,7 +8,7 @@ import pytest import yaml from cmdforge.tool import ( - Tool, ToolArgument, PromptStep, CodeStep, + Tool, ToolArgument, PromptStep, CodeStep, ToolStep, validate_tool_name, load_tool, save_tool, delete_tool, list_tools, tool_exists, create_wrapper_script, DEFAULT_CATEGORIES @@ -579,3 +579,40 @@ class TestDefaultCategories: assert "Developer" in DEFAULT_CATEGORIES assert "Data" in DEFAULT_CATEGORIES assert "Other" in DEFAULT_CATEGORIES + + +class TestAgentContext: + def test_toolstep_profile_roundtrip(self): + step = ToolStep(tool="my-tool", output_var="out", profile="architect") + d = step.to_dict() + assert d["profile"] == "architect" + restored = ToolStep.from_dict(d) + assert restored.profile == "architect" + + def test_toolstep_skills_roundtrip(self): + step = ToolStep(tool="my-tool", output_var="out", skills=["python"]) + d = step.to_dict() + assert d["skills"] == ["python"] + restored = ToolStep.from_dict(d) + assert restored.skills == ["python"] + + def test_toolstep_empty_skills_serialized(self): + step = ToolStep(tool="my-tool", output_var="out", skills=[]) + d = step.to_dict() + assert d["skills"] == [] + restored = ToolStep.from_dict(d) + assert restored.skills == [] + + def test_toolstep_no_skills_not_serialized(self): + step = ToolStep(tool="my-tool", output_var="out") + d = step.to_dict() + assert "skills" not in d + assert "profile" not in d + + def test_promptstep_skills_roundtrip(self): + step = PromptStep( + prompt="Test", provider="claude", output_var="out", + skills=["python"], + ) + restored = PromptStep.from_dict(step.to_dict()) + assert restored.skills == ["python"]