diff --git a/AGENTS.md b/AGENTS.md index 1d121a3..b7ed9ae 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -45,6 +45,8 @@ cf # Interactive tool picker - Use `cmdforge run-once "Instruction {input}"` with piped input for ad-hoc AI work; create a permanent tool only when the workflow is reusable. - Prefer CmdForge for local automation and composable workflows. Direct SDK/API integration is appropriate when the external API is a runtime dependency of the product itself. - Bootstrap supported coding hosts with `cmdforge mcp configure codex` or `cmdforge mcp configure claude-code`; preview changes with `--dry-run`. This updates only a marked policy block and never expands `server.expose`. +- Create project-owned tools with `cmdforge create NAME --project` or `forge-tool --name NAME --project`; do not edit CmdForge's source repository to add a consumer project's tools. +- For sensitive inputs, use `--no-fallback`, explicit provider constraints, and `--result-envelope json`; provenance is attached by CmdForge and must not be generated by the model. ## Testing diff --git a/CLAUDE.md b/CLAUDE.md index 7b8a7fb..4b4009a 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -45,6 +45,7 @@ python -m cmdforge.cli # Alternative CLI invocation - **resolver.py**: Tool resolution. `resolve_tool()` searches: project manifest → local tools → owner/name → registry. Returns `ResolvedTool` with path info - **collection.py**: Collection management (`Collection` dataclass, `resolve_tool_references()`, `classify_tool_reference()`), local collection storage in `~/.cmdforge/collections/` - **providers.py**: Provider abstraction. Supports subprocess CLI tools, OpenAI-compatible HTTP APIs, experimental PTY wrappers, fallback chains, and tool/MCP-server allowlists. Auto-discovers installed providers on first run. Config in `~/.cmdforge/providers.yaml` with versioned migration. +- **Provider execution contract**: Providers declare locality, typed capabilities, model identity, context/cost/latency hints, and an approved data classification. `ProviderExecutionPolicy` filters every fallback candidate; `ProviderResult` records runtime-owned provider/model provenance. - **skills.py**: Validated Agent Skills loader for per-provider `SKILL.md` context under `~/.cmdforge/providers//skills/` - **mcp_client.py / mcp_server.py**: Stdio and Streamable HTTP MCP integration, configuration, schema discovery, transport security, and exposure policy - **profiles.py**: AI persona profiles with system prompts, stored in `~/.cmdforge/profiles/` @@ -115,6 +116,11 @@ Use `cmdforge mcp configure codex` or `cmdforge mcp configure claude-code --scope project` to register the stdio MCP server and install a marked project-policy block. Preview with `--dry-run`. This bootstrap never adds tools to the MCP `server.expose` allowlist. +Project-owned tools belong under `./.cmdforge/`, created with +`cmdforge create NAME --project` or `forge-tool --name NAME --project`. Do not +edit CmdForge core merely to add another project's tools. Sensitive runs should +use `--no-fallback` and `--result-envelope json`; provenance must come from the +runtime rather than model-generated fields. ### Step Types diff --git a/README.md b/README.md index 417dc31..01fb0d3 100644 --- a/README.md +++ b/README.md @@ -124,6 +124,9 @@ Opens the graphical interface where you can create and manage tools visually. Fe - **Registry** - Search and install community tools from the CmdForge registry - **Providers** - Manage AI provider configurations - **Reuse & Discovery** - Compare similar tools and preview safe extraction of reusable steps +- **Scoped saving** - Save a new tool to the current project or your global tool collection + +![Tool Builder save-location selector](assets/screenshots/tool-builder-save-location.png) ### CLI Mode @@ -132,11 +135,15 @@ Opens the graphical interface where you can create and manage tools visually. Fe cmdforge list --filter summarize # Filter installed tools for humans cmdforge list --json --filter summarize --limit 10 # Compact agent catalog cmdforge list --json --full --limit 1 # Explicitly include prompts/code -cmdforge create mytool # Create new tool +cmdforge create mytool # Create a global tool +cmdforge create mytool --project # Versionable ./.cmdforge tool +echo "Build a bounded classifier" | forge-tool --name classify --project cmdforge edit mytool # Edit in $EDITOR cmdforge delete mytool # Delete tool cmdforge run mytool # Run a tool cat document.txt | cmdforge run-once "Summarize this:\n\n{input}" +cmdforge run private-tool --provider ollama --no-fallback \ + --require-local --data-classification private --result-envelope json cmdforge test mytool # Test with mock provider cmdforge check mytool # Check dependencies (meta-tools) cmdforge inspect mytool # Preflight, contract proposals, safe conformance tests diff --git a/assets/screenshots/tool-builder-save-location.png b/assets/screenshots/tool-builder-save-location.png new file mode 100644 index 0000000..8bc8669 Binary files /dev/null and b/assets/screenshots/tool-builder-save-location.png differ diff --git a/src/cmdforge/cli/__init__.py b/src/cmdforge/cli/__init__.py index 039a00a..741e039 100644 --- a/src/cmdforge/cli/__init__.py +++ b/src/cmdforge/cli/__init__.py @@ -2,6 +2,7 @@ import argparse import sys +from pathlib import Path from .. import __version__ @@ -50,6 +51,15 @@ def main(): p_create.add_argument("-p", "--prompt", help="Prompt template") p_create.add_argument("--provider", help="AI provider (default: mock)") p_create.add_argument("-f", "--force", action="store_true", help="Overwrite existing") + create_location = p_create.add_mutually_exclusive_group() + create_location.add_argument( + "--project", action="store_true", + help="Save under ./.cmdforge instead of ~/.cmdforge", + ) + create_location.add_argument( + "--output-dir", type=Path, + help="Save under an explicit tool root directory", + ) p_create.set_defaults(func=cmd_create) # 'edit' command @@ -81,6 +91,13 @@ def main(): p_run.add_argument("--show-prompt", action="store_true", help="Show prompts in addition to output") p_run.add_argument("-v", "--verbose", action="store_true", help="Show debug information") p_run.add_argument("--auto-install", action="store_true", help="Automatically install missing tool dependencies") + p_run.add_argument("--no-fallback", action="store_true", help="Fail instead of trying configured fallback providers") + p_run.add_argument("--require-local", action="store_true", help="Allow only providers declared local") + p_run.add_argument("--require-capability", action="append", default=[], help="Required provider capability (repeatable)") + p_run.add_argument("--data-classification", choices=["public", "internal", "private"], help="Input data classification") + p_run.add_argument("--require-model-identity", action="store_true") + p_run.add_argument("--require-model-digest", action="store_true") + p_run.add_argument("--result-envelope", choices=["json"], help="Wrap output with verified execution provenance") p_run.add_argument("tool_args", nargs=argparse.REMAINDER, help="Additional tool-specific arguments (use -- to separate)") p_run.set_defaults(func=cmd_run) @@ -100,6 +117,13 @@ def main(): ) p_run_once.add_argument("--max-tokens", type=int, help="Max output tokens") p_run_once.add_argument("--timeout", type=int, default=60, help="Timeout in seconds") + p_run_once.add_argument("--no-fallback", action="store_true") + p_run_once.add_argument("--require-local", action="store_true") + p_run_once.add_argument("--require-capability", action="append", default=[]) + p_run_once.add_argument("--data-classification", choices=["public", "internal", "private"]) + p_run_once.add_argument("--require-model-identity", action="store_true") + p_run_once.add_argument("--require-model-digest", action="store_true") + p_run_once.add_argument("--result-envelope", choices=["json"]) p_run_once.set_defaults(func=cmd_run_once) # 'ui' command (explicit) @@ -157,6 +181,7 @@ def main(): # providers list p_prov_list = providers_sub.add_parser("list", help="List all providers and their status") + p_prov_list.add_argument("--json", action="store_true") p_prov_list.set_defaults(func=cmd_providers) # providers check @@ -181,6 +206,27 @@ def main(): p_prov_add.add_argument("--model", help="Model ID for API providers") p_prov_add.add_argument("--api-key-env", help="Environment variable containing the API key") p_prov_add.add_argument("--tag", action="append", dest="tags", help="Provider tag (repeatable)") + p_prov_add.add_argument( + "--locality", choices=["local", "remote", "unknown"], + help="Where inference runs", + ) + p_prov_add.add_argument( + "--capability", action="append", dest="capabilities", + help="Typed capability such as structured-json (repeatable)", + ) + p_prov_add.add_argument("--model-digest", help="Verified model content digest") + p_prov_add.add_argument( + "--cost-class", choices=["free", "low", "medium", "high"] + ) + p_prov_add.add_argument( + "--latency-class", choices=["fast", "standard", "slow"] + ) + p_prov_add.add_argument("--max-context-tokens", type=int) + p_prov_add.add_argument( + "--data-policy", + choices=["unspecified", "public", "internal", "private"], + help="Maximum approved input classification", + ) p_prov_add.add_argument("--fallback", help="Fallback provider name; pass an empty value to clear") p_prov_add.add_argument( "--fallback-chain", @@ -667,9 +713,9 @@ def main(): def cmd_inspect(args): """Run preflight analysis on a tool.""" from ..preflight import analyze_tool - from ..tool import load_tool + from .tool_commands import _resolve_existing_tool - tool = load_tool(args.name) + tool = _resolve_existing_tool(args.name) if not tool: print(f"Error: Tool '{args.name}' not found.", file=sys.stderr) return 1 diff --git a/src/cmdforge/cli/provider_commands.py b/src/cmdforge/cli/provider_commands.py index ce4f86a..d725d10 100644 --- a/src/cmdforge/cli/provider_commands.py +++ b/src/cmdforge/cli/provider_commands.py @@ -353,6 +353,17 @@ def _cmd_providers_discover(args): def _cmd_providers_list(args): """List all providers and their status.""" providers = load_providers() + if getattr(args, "json", False): + import json + payload = [] + for provider in providers: + available, status = _provider_status(provider) + item = provider.to_dict() + item["available"] = available + item["status"] = status + payload.append(item) + print(json.dumps(payload, indent=2)) + return 0 print(f"Configured providers ({len(providers)}):\n") for p in providers: exists, status = _provider_status(p) @@ -365,6 +376,15 @@ def _cmd_providers_list(args): print(f" Status: {status}") if p.description: print(f" Info: {p.description}") + print(f" Locality: {p.locality}") + print(f" Data policy: {p.data_policy}") + if p.capabilities: + print(f" Capabilities: {', '.join(p.capabilities)}") + if p.model: + identity = p.model + if p.model_digest: + identity += f" ({p.model_digest})" + print(f" Model: {identity}") print() return 0 @@ -379,6 +399,17 @@ def _cmd_providers_add(args): model = args.model if args.model is not None else (existing.model if existing else None) api_key_env = args.api_key_env if args.api_key_env is not None else (existing.api_key_env if existing else None) tags = args.tags if args.tags is not None else (existing.tags if existing else []) + locality = args.locality if args.locality is not None else (existing.locality if existing else "unknown") + capabilities = args.capabilities if args.capabilities is not None else (existing.capabilities if existing else []) + model_digest = args.model_digest if args.model_digest is not None else (existing.model_digest if existing else None) + cost_class = args.cost_class if args.cost_class is not None else (existing.cost_class if existing else None) + latency_class = args.latency_class if args.latency_class is not None else (existing.latency_class if existing else None) + max_context_tokens = ( + args.max_context_tokens + if args.max_context_tokens is not None + else existing.max_context_tokens if existing else None + ) + data_policy = args.data_policy if args.data_policy is not None else (existing.data_policy if existing else "unspecified") fallback = args.fallback if args.fallback is not None else (existing.fallback if existing else None) fallback = fallback or None @@ -406,6 +437,15 @@ def _cmd_providers_add(args): fallback_chain=fallback_chain, api_key_env=api_key_env, pty_config=existing.pty_config if existing else None, + tools=existing.tools if existing else None, + mcp_servers=existing.mcp_servers if existing else None, + locality=locality, + capabilities=capabilities, + model_digest=model_digest, + cost_class=cost_class, + latency_class=latency_class, + max_context_tokens=max_context_tokens, + data_policy=data_policy, ) add_provider(provider) print(f"Provider '{name}' added/updated.") diff --git a/src/cmdforge/cli/tool_commands.py b/src/cmdforge/cli/tool_commands.py index be5c1e0..8c2ad35 100644 --- a/src/cmdforge/cli/tool_commands.py +++ b/src/cmdforge/cli/tool_commands.py @@ -5,11 +5,105 @@ from pathlib import Path from ..tool import ( list_tools, load_tool, save_tool, delete_tool, get_tools_dir, - Tool, ToolArgument, PromptStep, CodeStep, ToolStep, validate_tool_name + get_project_tools_dir, Tool, ToolArgument, PromptStep, CodeStep, ToolStep, + validate_tool_name ) from ..gui import run_gui +def _resolve_existing_tool(name: str): + """Resolve project-local before global without registry auto-fetch.""" + from ..resolver import find_tool + + project_candidate = any( + (Path.cwd() / root / name / "config.yaml").exists() + for root in (".cmdforge", "cmdforge") + ) + if project_candidate: + resolved = find_tool(name) + if resolved: + return resolved.tool + from .. import tool as tool_module + return tool_module.load_tool(name) + + +def _project_tool_names(project_dir=None) -> list[str]: + """List project-local references using the same two resolver roots.""" + project_dir = (project_dir or Path.cwd()).resolve() + names = [] + for root_name in (".cmdforge", "cmdforge"): + root = project_dir / root_name + if not root.is_dir(): + continue + for item in sorted(root.iterdir()): + if not item.is_dir() or item.name.startswith("."): + continue + if (item / "config.yaml").is_file(): + names.append(item.name) + continue + for child in sorted(item.iterdir()): + if child.is_dir() and (child / "config.yaml").is_file(): + names.append(f"{item.name}/{child.name}") + return list(dict.fromkeys(names)) + + +def _execution_policy_from_args(*namespaces): + """Combine top-level and wrapper-compatible execution policy flags.""" + from ..providers import ProviderExecutionPolicy + + capabilities = [] + for namespace in namespaces: + capabilities.extend( + getattr(namespace, "require_capability", None) or [] + ) + data_classification = next(( + getattr(namespace, "data_classification", None) + for namespace in namespaces + if getattr(namespace, "data_classification", None) + ), "public") + return ProviderExecutionPolicy( + allow_fallback=not any( + getattr(namespace, "no_fallback", False) for namespace in namespaces + ), + required_locality=( + "local" if any( + getattr(namespace, "require_local", False) + for namespace in namespaces + ) else None + ), + required_capabilities=tuple(dict.fromkeys(capabilities)), + data_classification=data_classification, + require_model_identity=any( + getattr(namespace, "require_model_identity", False) + for namespace in namespaces + ), + require_model_digest=any( + getattr(namespace, "require_model_digest", False) + for namespace in namespaces + ), + ) + + +def _result_envelope( + output: str, provider_calls: list[dict], exit_code: int = 0 +) -> str: + """Build a stable machine-readable envelope without parsing tool output.""" + import json + + execution = { + "exit_code": exit_code, + "provider_calls": provider_calls, + "fallback_used": any( + call.get("fallback_used", False) for call in provider_calls + ), + "actual_providers": list(dict.fromkeys( + call["actual_provider"] for call in provider_calls + if call.get("actual_provider") + )), + } + return json.dumps({"output": output, "execution": execution}, indent=2) + + def cmd_list(args): """List all tools.""" import json @@ -24,8 +118,10 @@ def cmd_list(args): return 1 matches = [] - for name in sorted(list_tools()): - tool = load_tool(name) + project_names = _project_tool_names() + all_names = list(dict.fromkeys([*project_names, *sorted(list_tools())])) + for name in all_names: + tool = _resolve_existing_tool(name) if not tool: continue if filter_text and ( @@ -151,15 +247,22 @@ def cmd_run_once(args): return 1 print(f"[run-once] Provider: {provider}", file=sys.stderr) + policy = _execution_policy_from_args(args) result = call_provider( - provider, prompt, timeout=timeout, max_tokens=max_tokens + provider, prompt, timeout=timeout, max_tokens=max_tokens, + execution_policy=policy, ) if not result.success: + if getattr(args, "result_envelope", None) == "json": + print(_result_envelope("", [result.provenance()], exit_code=1)) print(f"Error: {result.error}", file=sys.stderr) return 1 - print(result.text) + if getattr(args, "result_envelope", None) == "json": + print(_result_envelope(result.text, [result.provenance()])) + else: + print(result.text) return 0 @@ -242,10 +345,22 @@ def cmd_create(args): print(f"Error: Invalid tool name '{name}': {error_msg}") return 1 - # Check if already exists - existing = load_tool(name) - if existing and not args.force: - print(f"Error: Tool '{name}' already exists. Use --force to overwrite.") + output_dir = getattr(args, "output_dir", None) + project = getattr(args, "project", False) + target_root = ( + Path(output_dir).resolve() + if output_dir is not None + else get_project_tools_dir() if project else get_tools_dir() + ) + + # Check only the requested scope. A project tool may intentionally shadow + # a global tool with the same name. + config_path = target_root / name / "config.yaml" + if config_path.exists() and not args.force: + print( + f"Error: Tool '{name}' already exists at {config_path}. " + "Use --force to overwrite." + ) return 1 # Create a tool with a single prompt step @@ -266,11 +381,19 @@ def cmd_create(args): output="{response}" if steps else "{input}" ) - path = save_tool(tool) + path = save_tool(tool, tools_dir=target_root) print(f"Created tool '{name}'") print(f"Config: {path}") - print(f"\nUse 'cmdforge ui' to add arguments, steps, and customize.") - print(f"Or run: {name} < input.txt") + print("\nUse 'cmdforge ui' to add arguments, steps, and customize.") + if project: + print(f"Run from its project: cmdforge run {name} < input.txt") + elif output_dir is not None: + print( + "Saved under the requested tool root. CmdForge automatically " + "discovers roots named .cmdforge or cmdforge." + ) + else: + print(f"Or run: {name} < input.txt") return 0 @@ -323,7 +446,7 @@ def cmd_delete(args): def cmd_test(args): """Test a tool with mock provider.""" - tool = load_tool(args.name) + tool = _resolve_existing_tool(args.name) if not tool: print(f"Error: Tool '{args.name}' not found.") return 1 @@ -360,7 +483,7 @@ def cmd_run(args): """Run a tool.""" from ..runner import collect_custom_args, create_argument_parser, run_tool - tool = load_tool(args.name) + tool = _resolve_existing_tool(args.name) if not tool: print(f"Error: Tool '{args.name}' not found.", file=sys.stderr) return 1 @@ -404,6 +527,12 @@ def cmd_run(args): show_prompt = args.show_prompt or parsed_tool_args.show_prompt verbose = args.verbose or parsed_tool_args.verbose auto_install = args.auto_install or parsed_tool_args.auto_install + execution_policy = _execution_policy_from_args(args, parsed_tool_args) + result_envelope = ( + getattr(args, "result_envelope", None) + or getattr(parsed_tool_args, "result_envelope", None) + ) + provider_calls = [] # Run tool output, code = run_tool( @@ -414,16 +543,22 @@ def cmd_run(args): dry_run=dry_run, show_prompt=show_prompt, verbose=verbose, - auto_install=auto_install + auto_install=auto_install, + execution_policy=execution_policy, + provenance_log=provider_calls, ) # Write output - if code == 0 and output: + if result_envelope == "json" or (code == 0 and output): + rendered_output = ( + _result_envelope(output, provider_calls, exit_code=code) + if result_envelope == "json" else output + ) output_file = args.output or parsed_tool_args.output_file if output_file: - Path(output_file).write_text(output) + Path(output_file).write_text(rendered_output) else: - print(output) + print(rendered_output) return code diff --git a/src/cmdforge/gui/dialogs/provider_dialog.py b/src/cmdforge/gui/dialogs/provider_dialog.py index 3f4e477..50abf59 100644 --- a/src/cmdforge/gui/dialogs/provider_dialog.py +++ b/src/cmdforge/gui/dialogs/provider_dialog.py @@ -44,6 +44,44 @@ class ProviderDialog(QDialog): self.desc_input.setPlaceholderText("Claude AI via claude-cli") form.addRow("Description:", self.desc_input) + self.locality_combo = QComboBox() + for value in ("unknown", "local", "remote"): + self.locality_combo.addItem(value.title(), value) + form.addRow("Locality:", self.locality_combo) + + self.data_policy_combo = QComboBox() + for value in ("unspecified", "public", "internal", "private"): + self.data_policy_combo.addItem(value.title(), value) + form.addRow("Data policy:", self.data_policy_combo) + + self.capabilities_input = QLineEdit() + self.capabilities_input.setPlaceholderText("text, structured-json, reasoning") + form.addRow("Capabilities:", self.capabilities_input) + + self.model_input = QLineEdit() + self.model_input.setPlaceholderText("Model ID, if known") + form.addRow("Model:", self.model_input) + + self.model_digest_input = QLineEdit() + self.model_digest_input.setPlaceholderText("sha256:... (verified local models)") + form.addRow("Model digest:", self.model_digest_input) + + self.max_context_input = QLineEdit() + self.max_context_input.setPlaceholderText("Maximum context tokens") + form.addRow("Context tokens:", self.max_context_input) + + self.cost_class_combo = QComboBox() + self.cost_class_combo.addItem("Unspecified", None) + for value in ("free", "low", "medium", "high"): + self.cost_class_combo.addItem(value.title(), value) + form.addRow("Cost class:", self.cost_class_combo) + + self.latency_class_combo = QComboBox() + self.latency_class_combo.addItem("Unspecified", None) + for value in ("fast", "standard", "slow"): + self.latency_class_combo.addItem(value.title(), value) + form.addRow("Latency class:", self.latency_class_combo) + self.fallback_combo = QComboBox() self.fallback_combo.addItem("(none)", None) # Populate with available providers @@ -86,6 +124,24 @@ class ProviderDialog(QDialog): self.name_input.setEnabled(False) # Can't rename self.cmd_input.setText(provider.command) self.desc_input.setText(provider.description or "") + self.locality_combo.setCurrentIndex( + max(self.locality_combo.findData(provider.locality), 0) + ) + self.data_policy_combo.setCurrentIndex( + max(self.data_policy_combo.findData(provider.data_policy), 0) + ) + self.capabilities_input.setText(", ".join(provider.capabilities)) + self.model_input.setText(provider.model or "") + self.model_digest_input.setText(provider.model_digest or "") + self.max_context_input.setText( + str(provider.max_context_tokens) if provider.max_context_tokens else "" + ) + self.cost_class_combo.setCurrentIndex( + max(self.cost_class_combo.findData(provider.cost_class), 0) + ) + self.latency_class_combo.setCurrentIndex( + max(self.latency_class_combo.findData(provider.latency_class), 0) + ) # Set fallback selection if provider.fallback: idx = self.fallback_combo.findData(provider.fallback) @@ -107,6 +163,21 @@ class ProviderDialog(QDialog): description = self.desc_input.text().strip() fallback = self.fallback_combo.currentData() + capabilities = [ + value.strip() for value in self.capabilities_input.text().split(",") + if value.strip() + ] + try: + max_context_tokens = ( + int(self.max_context_input.text().strip()) + if self.max_context_input.text().strip() else None + ) + except ValueError: + from PySide6.QtWidgets import QMessageBox + QMessageBox.warning( + self, "Invalid", "Context tokens must be a positive integer." + ) + return # Prevent self-referential fallback if fallback == name: @@ -122,12 +193,21 @@ class ProviderDialog(QDialog): description=description, fallback=fallback, type=existing.type if existing else "subprocess", - model=existing.model if existing else None, + model=self.model_input.text().strip() or None, tags=existing.tags if existing else [], install=existing.install if existing else None, fallback_chain=existing.fallback_chain if existing else None, api_key_env=existing.api_key_env if existing else None, pty_config=existing.pty_config if existing else None, + tools=existing.tools if existing else None, + mcp_servers=existing.mcp_servers if existing else None, + locality=self.locality_combo.currentData(), + capabilities=capabilities, + model_digest=self.model_digest_input.text().strip() or None, + cost_class=self.cost_class_combo.currentData(), + latency_class=self.latency_class_combo.currentData(), + max_context_tokens=max_context_tokens, + data_policy=self.data_policy_combo.currentData(), )) self.accept() except Exception as e: diff --git a/src/cmdforge/gui/dialogs/step_dialog.py b/src/cmdforge/gui/dialogs/step_dialog.py index 36bd35d..ec3b465 100644 --- a/src/cmdforge/gui/dialogs/step_dialog.py +++ b/src/cmdforge/gui/dialogs/step_dialog.py @@ -321,6 +321,15 @@ class PromptStepDialog(QDialog): self.strip_fences_check = QCheckBox("Strip markdown code fences from output") form.addRow("", self.strip_fences_check) + self.allow_fallback_check = QCheckBox( + "Allow configured fallback providers if this provider fails" + ) + self.allow_fallback_check.setChecked(True) + self.allow_fallback_check.setToolTip( + "Disable for privacy-sensitive or reproducible steps that must fail closed." + ) + form.addRow("", self.allow_fallback_check) + # Structured output options form.addRow(QLabel("")) # Spacer structured_label = QLabel("Structured Output") @@ -455,6 +464,7 @@ class PromptStepDialog(QDialog): self.output_input.setText(step.output_var) self.prompt_input.setPlainText(step.prompt) self.strip_fences_check.setChecked(step.strip_fences) + self.allow_fallback_check.setChecked(step.fallback_policy != "deny") # Structured output fields self.plain_text_check.setChecked(step.plain_text) @@ -497,7 +507,10 @@ class PromptStepDialog(QDialog): strip_fences=self.strip_fences_check.isChecked(), output_schema=self._output_schema, plain_text=self.plain_text_check.isChecked(), - max_retries=self.retries_spin.value() + max_retries=self.retries_spin.value(), + fallback_policy=( + "allow" if self.allow_fallback_check.isChecked() else "deny" + ), ) diff --git a/src/cmdforge/gui/pages/tool_builder_page.py b/src/cmdforge/gui/pages/tool_builder_page.py index f594fe0..fbe5650 100644 --- a/src/cmdforge/gui/pages/tool_builder_page.py +++ b/src/cmdforge/gui/pages/tool_builder_page.py @@ -3,6 +3,7 @@ import copy import difflib import re +from pathlib import Path import yaml @@ -17,8 +18,8 @@ from PySide6.QtCore import Qt, QThread, Signal, QTimer from ...tool import ( Tool, ToolArgument, PromptStep, CodeStep, ToolStep, - load_tool, save_tool, tool_exists, validate_tool_name, get_all_categories, - ensure_settings + load_tool, save_tool, validate_tool_name, get_all_categories, + ensure_settings, get_project_tools_dir, get_tools_dir ) from ..widgets.icons import get_prompt_icon, get_code_icon, get_tool_icon @@ -107,6 +108,21 @@ def _switch_to_existing_tool(main_window, name: str) -> None: main_window.open_tool_builder(name) +def default_save_scope(project_dir: Path) -> str: + """Prefer project scope only when the cwd is recognizably a project.""" + project_dir = Path(project_dir) + if any((project_dir / marker).exists() for marker in ( + "cmdforge.yaml", ".cmdforge", ".git", "pyproject.toml", "package.json", + )): + return "project" + return "global" + + +def existing_tool_root(config_path: Path) -> Path: + """Return the tool-root directory containing an existing named tool.""" + return Path(config_path).resolve().parent.parent + + class ToolBuilderPage(QWidget): """Tool builder/editor page.""" @@ -185,6 +201,23 @@ class ToolBuilderPage(QWidget): self.category_combo.addItem(cat) info_layout.addRow("Category:", self.category_combo) + self.save_location_combo = QComboBox() + project_root = get_project_tools_dir() + self.save_location_combo.addItem( + f"This project ({project_root})", "project" + ) + self.save_location_combo.addItem( + f"My tools ({get_tools_dir()})", "global" + ) + preferred_scope = default_save_scope(Path.cwd()) + preferred_index = self.save_location_combo.findData(preferred_scope) + self.save_location_combo.setCurrentIndex(max(preferred_index, 0)) + self.save_location_combo.setToolTip( + "Project tools are versionable and resolve only from this project; " + "My tools are available globally and receive a shell wrapper." + ) + info_layout.addRow("Save tool:", self.save_location_combo) + left_layout.addWidget(info_box) # Reuse guidance is advisory and never changes a draft automatically. @@ -658,7 +691,8 @@ class ToolBuilderPage(QWidget): if not valid: QMessageBox.warning(self, "Validation", error) return - if tool_exists(name): + extraction_root = self._selected_tools_dir() + if (extraction_root / name / "config.yaml").exists(): QMessageBox.warning(self, "Tool Already Exists", f"'{name}' already exists.") return try: @@ -684,7 +718,7 @@ class ToolBuilderPage(QWidget): if preview.clickedButton() is not create_button: return try: - save_tool(extracted) + save_tool(extracted, tools_dir=extraction_root) except (OSError, ValueError) as exc: QMessageBox.critical(self, "Extraction Failed", str(exc)) return @@ -807,6 +841,17 @@ class ToolBuilderPage(QWidget): self._tool = tool self.name_input.setText(tool.name) self.name_input.setEnabled(False) # Can't rename + if tool.path: + project_root = get_project_tools_dir() + scope = ( + "project" + if existing_tool_root(tool.path) == project_root.resolve() + else "global" + ) + index = self.save_location_combo.findData(scope) + if index >= 0: + self.save_location_combo.setCurrentIndex(index) + self.save_location_combo.setEnabled(False) self.desc_input.setText(tool.description or "") # Set category @@ -1371,9 +1416,13 @@ class ToolBuilderPage(QWidget): QMessageBox.warning(self, "Validation", error) return - # Block silent overwrite when creating a new tool - if (not self.editing or name != self.original_name) and tool_exists(name): - existing = load_tool(name) + target_root = self._selected_tools_dir() + + # Block silent overwrite in the selected scope. A project-local tool + # may intentionally shadow a global tool with the same name. + target_config = target_root / name / "config.yaml" + if (not self.editing or name != self.original_name) and target_config.exists(): + existing = load_tool(name) if self.save_location_combo.currentData() == "global" else None existing_desc = f" ({existing.description})" if existing and existing.description else "" msg = QMessageBox(self) msg.setIcon(QMessageBox.Warning) @@ -1393,7 +1442,7 @@ class ToolBuilderPage(QWidget): return elif clicked is btn_copy: suffix = 2 - while tool_exists(f"{name}-{suffix}"): + while (target_root / f"{name}-{suffix}" / "config.yaml").exists(): suffix += 1 copy_name = f"{name}-{suffix}" self.name_input.setText(copy_name) @@ -1514,7 +1563,7 @@ class ToolBuilderPage(QWidget): tool.version = self._tool.version try: - config_path = save_tool(tool) + config_path = save_tool(tool, tools_dir=target_root) tool_dir = config_path.parent # Save defaults if provided @@ -1551,6 +1600,14 @@ class ToolBuilderPage(QWidget): """Cancel and return to tools page.""" self.main_window.close_tool_builder() + def _selected_tools_dir(self) -> Path: + """Return the save root selected by the author.""" + if self.editing and self._tool and self._tool.path: + return existing_tool_root(self._tool.path) + if self.save_location_combo.currentData() == "project": + return get_project_tools_dir() + return get_tools_dir().resolve() + def save_tool(self): """Public method for keyboard shortcut to save the tool.""" self._save() diff --git a/src/cmdforge/providers.py b/src/cmdforge/providers.py index 0ec457c..a73f50a 100644 --- a/src/cmdforge/providers.py +++ b/src/cmdforge/providers.py @@ -55,6 +55,9 @@ KNOWN_PROVIDER_CLIS = { "command": "opencode run --model opencode/big-pickle", "description": "OpenCode - Big Pickle (free general model)", "tags": ["free", "code", "general"], + "locality": "remote", + "capabilities": ["text", "structured-json", "reasoning"], + "data_policy": "public", "install_group": "opencode", }, "agy": { @@ -62,6 +65,9 @@ KNOWN_PROVIDER_CLIS = { "command": "agy -p", "description": "Antigravity - Google free tier, Gemini models", "tags": ["free-tier", "code", "large-context"], + "locality": "remote", + "capabilities": ["text", "structured-json", "reasoning"], + "data_policy": "public", "install_group": "agy", }, "codex": { @@ -69,6 +75,9 @@ KNOWN_PROVIDER_CLIS = { "command": "codex exec -", "description": "Codex CLI - OpenAI free tier available", "tags": ["free-tier", "code", "general"], + "locality": "remote", + "capabilities": ["text", "structured-json", "reasoning"], + "data_policy": "public", "install_group": "codex", }, "claude": { @@ -76,6 +85,9 @@ KNOWN_PROVIDER_CLIS = { "command": "claude -p", "description": "Claude Code - auto-routes to best model", "tags": ["paid", "subscription", "code"], + "locality": "remote", + "capabilities": ["text", "structured-json", "reasoning"], + "data_policy": "public", "install_group": "claude", }, "crush": { @@ -83,6 +95,9 @@ KNOWN_PROVIDER_CLIS = { "command": "crush run --quiet", "description": "Crush - multi-model via Hyper credits or API keys", "tags": ["free-tier", "multi", "code"], + "locality": "remote", + "capabilities": ["text", "structured-json", "reasoning"], + "data_policy": "public", "install_group": "crush", }, "ollama": { @@ -90,6 +105,10 @@ KNOWN_PROVIDER_CLIS = { "command": "ollama run llama3.2", "description": "Ollama - local, private, free", "tags": ["free", "local", "private"], + "model": "llama3.2", + "locality": "local", + "capabilities": ["text", "structured-json"], + "data_policy": "private", "install_group": "ollama", }, } @@ -102,6 +121,9 @@ KNOWN_API_KEYS = { "model": "openrouter/auto-beta", "description": "OpenRouter - 300+ models, one API key", "tags": ["api", "per-token", "multi"], + "locality": "remote", + "capabilities": ["text", "structured-json", "reasoning"], + "data_policy": "public", }, "DEEPSEEK_API_KEY": { "name": "deepseek-api", @@ -109,6 +131,9 @@ KNOWN_API_KEYS = { "model": "deepseek-chat", "description": "DeepSeek API - inexpensive per-token access", "tags": ["api", "per-token", "cheap"], + "locality": "remote", + "capabilities": ["text", "structured-json", "reasoning"], + "data_policy": "public", }, "OPENAI_API_KEY": { "name": "openai-api", @@ -116,6 +141,9 @@ KNOWN_API_KEYS = { "model": "gpt-4o", "description": "OpenAI API - direct GPT access", "tags": ["api", "per-token", "code"], + "locality": "remote", + "capabilities": ["text", "structured-json", "reasoning"], + "data_policy": "public", }, } @@ -181,6 +209,10 @@ def discover_installed_providers() -> List[dict]: "command": f"ollama run {model_name}", "description": f"Ollama local model: {model_name}", "tags": ["free", "local", "private"], + "model": model_name, + "locality": "local", + "capabilities": ["text", "structured-json"], + "data_policy": "private", }) except (subprocess.TimeoutExpired, FileNotFoundError): pass @@ -188,6 +220,26 @@ def discover_installed_providers() -> List[dict]: return found +def _provider_from_discovery(data: dict) -> "Provider": + """Convert trusted discovery facts without persisting scanner-only keys.""" + return Provider( + name=data["name"], + command=data["command"], + description=data.get("description", ""), + type=data.get("type", "subprocess"), + model=data.get("model"), + tags=data.get("tags", []), + api_key_env=data.get("env_var"), + locality=data.get("locality", "unknown"), + capabilities=data.get("capabilities", []), + model_digest=data.get("model_digest"), + cost_class=data.get("cost_class"), + latency_class=data.get("latency_class"), + max_context_tokens=data.get("max_context_tokens"), + data_policy=data.get("data_policy", "unspecified"), + ) + + @dataclass class Provider: """Definition of an AI provider. @@ -215,6 +267,13 @@ class Provider: 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 + locality: str = "unknown" # "local" | "remote" | "unknown" + capabilities: List[str] = field(default_factory=list) + model_digest: Optional[str] = None + cost_class: Optional[str] = None # "free" | "low" | "medium" | "high" + latency_class: Optional[str] = None # "fast" | "standard" | "slow" + max_context_tokens: Optional[int] = None + data_policy: str = "unspecified" # max: public | internal | private def __post_init__(self) -> None: for field_name, values in ( @@ -228,6 +287,30 @@ class Provider: raise ValueError( f"Provider {field_name} must be a list of non-empty strings or null" ) + if self.locality not in {"local", "remote", "unknown"}: + raise ValueError("Provider locality must be local, remote, or unknown") + if ( + not isinstance(self.capabilities, list) + or not all(isinstance(value, str) and value for value in self.capabilities) + ): + raise ValueError("Provider capabilities must be a list of non-empty strings") + if self.cost_class not in {None, "free", "low", "medium", "high"}: + raise ValueError("Provider cost_class must be free, low, medium, or high") + if self.latency_class not in {None, "fast", "standard", "slow"}: + raise ValueError("Provider latency_class must be fast, standard, or slow") + if self.data_policy not in {"unspecified", "public", "internal", "private"}: + raise ValueError( + "Provider data_policy must be unspecified, public, internal, or private" + ) + if ( + self.max_context_tokens is not None + and ( + not isinstance(self.max_context_tokens, int) + or isinstance(self.max_context_tokens, bool) + or self.max_context_tokens <= 0 + ) + ): + raise ValueError("Provider max_context_tokens must be a positive integer") def to_dict(self) -> dict: d = { @@ -256,6 +339,20 @@ class Provider: d["tools"] = self.tools if self.mcp_servers is not None: d["mcp_servers"] = self.mcp_servers + if self.locality != "unknown": + d["locality"] = self.locality + if self.capabilities: + d["capabilities"] = self.capabilities + if self.model_digest: + d["model_digest"] = self.model_digest + if self.cost_class: + d["cost_class"] = self.cost_class + if self.latency_class: + d["latency_class"] = self.latency_class + if self.max_context_tokens is not None: + d["max_context_tokens"] = self.max_context_tokens + if self.data_policy != "unspecified": + d["data_policy"] = self.data_policy return d @classmethod @@ -265,7 +362,7 @@ class Provider: command=data["command"], description=data.get("description", ""), type=data.get("type", "subprocess"), - model=data.get("model"), + model=data.get("model", data.get("model_id")), fallback=data.get("fallback"), tags=data.get("tags", []) or [], install=data.get("install"), @@ -274,6 +371,15 @@ class Provider: pty_config=data.get("pty_config"), tools=data.get("tools"), mcp_servers=data.get("mcp_servers"), + locality=data.get("locality", "unknown"), + capabilities=data.get("capabilities", []) or [], + model_digest=data.get("model_digest"), + cost_class=data.get("cost_class"), + latency_class=data.get("latency_class"), + max_context_tokens=data.get( + "max_context_tokens", data.get("max_input_tokens") + ), + data_policy=data.get("data_policy", "unspecified"), ) @@ -283,6 +389,164 @@ class ProviderResult: text: str success: bool error: Optional[str] = None + requested_provider: Optional[str] = None + actual_provider: Optional[str] = None + model: Optional[str] = None + model_digest: Optional[str] = None + attempted_providers: List[str] = field(default_factory=list) + fallback_used: bool = False + locality: str = "unknown" + model_identity_source: str = "unknown" + + def provenance(self) -> dict: + """Return runtime-owned execution facts safe for result envelopes.""" + data = { + "success": self.success, + "requested_provider": self.requested_provider, + "actual_provider": self.actual_provider, + "attempted_providers": list(self.attempted_providers), + "fallback_used": self.fallback_used, + "model": self.model, + "model_digest": self.model_digest, + "locality": self.locality, + "model_identity_source": self.model_identity_source, + } + if self.error: + data["error"] = self.error + return data + + +@dataclass(frozen=True) +class ProviderExecutionPolicy: + """Caller-supplied, fail-closed constraints for provider execution.""" + + allow_fallback: bool = True + required_locality: Optional[str] = None + required_capabilities: tuple[str, ...] = () + data_classification: str = "public" + require_model_identity: bool = False + require_model_digest: bool = False + + def __post_init__(self) -> None: + if self.required_locality not in {None, "local", "remote"}: + raise ValueError("required_locality must be local, remote, or null") + if self.data_classification not in {"public", "internal", "private"}: + raise ValueError( + "data_classification must be public, internal, or private" + ) + + +_DATA_CLASS_LEVEL = {"public": 0, "internal": 1, "private": 2} + + +def effective_provider_locality(provider: Provider) -> str: + """Account for Ollama clients redirected to a non-loopback host.""" + try: + executable = shlex.split(os.path.expandvars(provider.command))[0] + except (IndexError, ValueError): + executable = "" + if Path(executable).name != "ollama": + return provider.locality + ollama_host = os.environ.get("OLLAMA_HOST", "").strip() + if not ollama_host: + return provider.locality + from urllib.parse import urlparse + parsed = urlparse( + ollama_host if "://" in ollama_host else f"//{ollama_host}" + ) + hostname = (parsed.hostname or "").lower() + if hostname in {"localhost", "127.0.0.1", "::1"}: + return "local" + return "remote" + + +def runtime_provider_identity(provider: Provider) -> tuple[Optional[str], Optional[str], str]: + """Return observed model identity when supported, else configured facts.""" + configured = ( + provider.model, + provider.model_digest, + "provider-config" if provider.model or provider.model_digest else "unknown", + ) + try: + parts = shlex.split(os.path.expandvars(provider.command)) + except ValueError: + return configured + if not parts or Path(parts[0]).name != "ollama": + return configured + if effective_provider_locality(provider) != "local": + return configured + try: + run_index = parts.index("run") + model = provider.model or parts[run_index + 1] + except (ValueError, IndexError): + return configured + + host = os.environ.get("OLLAMA_HOST", "http://127.0.0.1:11434").strip() + if "://" not in host: + host = f"http://{host}" + from urllib.request import Request, urlopen + import json + + request = Request(host.rstrip("/") + "/api/tags") + api_key = os.environ.get("OLLAMA_API_KEY") + if api_key: + request.add_header("Authorization", f"Bearer {api_key}") + try: + with urlopen(request, timeout=2) as response: + payload = json.loads(response.read().decode("utf-8")) + except Exception: + return configured + wanted = {model, f"{model}:latest"} + for item in payload.get("models", []): + names = {item.get("name"), item.get("model")} + if names & wanted or any( + value and value.removesuffix(":latest") == model.removesuffix(":latest") + for value in names + ): + digest = item.get("digest") or provider.model_digest + if digest and not digest.startswith("sha256:"): + digest = f"sha256:{digest}" + return item.get("model") or item.get("name") or model, digest, "ollama-api" + return configured + + +def provider_policy_error( + provider: Provider, + policy: ProviderExecutionPolicy, + runtime_identity: Optional[tuple[Optional[str], Optional[str], str]] = None, +) -> Optional[str]: + """Return why a provider is ineligible, or None when it is allowed.""" + if ( + policy.required_locality is not None + and effective_provider_locality(provider) != policy.required_locality + ): + return ( + f"Provider '{provider.name}' locality is " + f"{effective_provider_locality(provider)}; " + f"{policy.required_locality} is required" + ) + missing = sorted(set(policy.required_capabilities) - set(provider.capabilities)) + if missing: + return ( + f"Provider '{provider.name}' lacks required capabilities: " + + ", ".join(missing) + ) + if policy.data_classification != "public": + allowed = _DATA_CLASS_LEVEL.get(provider.data_policy, -1) + needed = _DATA_CLASS_LEVEL[policy.data_classification] + if allowed < needed: + return ( + f"Provider '{provider.name}' data_policy is {provider.data_policy}; " + f"{policy.data_classification} data requires an explicit compatible policy" + ) + model, digest, _ = runtime_identity or ( + provider.model, provider.model_digest, "provider-config" + ) + if policy.require_model_identity and not model: + return f"Provider '{provider.name}' does not declare a model identity" + if policy.require_model_digest and not digest: + return f"Provider '{provider.name}' does not declare a model digest" + return None # Default providers that come pre-configured @@ -334,7 +598,12 @@ DEFAULT_PROVIDERS = [ # LOCAL MODELS Provider("ollama", "ollama run llama3.2", "Ollama - local, private, free (GPU recommended)", - tags=["free", "local", "private"]), + model="llama3.2", + tags=["free", "local", "private"], + locality="local", + capabilities=["text", "structured-json"], + cost_class="free", + data_policy="private"), # API-TYPE PROVIDERS (pay-per-token, fallback when no CLI covers model) Provider("openrouter", @@ -343,18 +612,26 @@ DEFAULT_PROVIDERS = [ type="api", model="openrouter/auto-beta", api_key_env="OPENROUTER_API_KEY", - tags=["api", "per-token", "multi", "fallback"]), + tags=["api", "per-token", "multi", "fallback"], + locality="remote", + capabilities=["text", "structured-json", "reasoning"], + data_policy="public"), Provider("deepseek-api", "https://api.deepseek.com/v1", "DeepSeek API - inexpensive per-token access", type="api", model="deepseek-chat", api_key_env="DEEPSEEK_API_KEY", - tags=["api", "per-token", "cheap", "reasoning"]), + tags=["api", "per-token", "cheap", "reasoning"], + locality="remote", + capabilities=["text", "structured-json", "reasoning"], + data_policy="public"), # Mock for testing Provider("mock", "mock", "Mock provider for testing", - tags=["testing"]), + model="mock", tags=["testing"], locality="local", + capabilities=["text", "structured-json"], + data_policy="private"), ] @@ -379,10 +656,7 @@ def get_providers_file() -> Path: for d in cli_found: name = d["name"] desc = d.get("description", "") - selected.append(Provider( - name=name, command=d["command"], description=desc, - tags=d.get("tags", []), - )) + selected.append(_provider_from_discovery(d)) print(f" [+] {name:20s} {desc}", file=sys.stderr) if api_found: @@ -392,11 +666,7 @@ def get_providers_file() -> Path: desc = d.get("description", "") model = d.get("model", "auto") env_var = d.get("env_var", "") - selected.append(Provider( - name=name, command=d["command"], description=desc, - type="api", model=model, api_key_env=env_var, - tags=d.get("tags", []), - )) + selected.append(_provider_from_discovery(d)) print(f" [+] {name:20s} ({model})", file=sys.stderr) if ollama_found: @@ -404,10 +674,7 @@ def get_providers_file() -> Path: for d in ollama_found[:10]: name = d["name"] desc = d.get("description", "") - selected.append(Provider( - name=name, command=d["command"], description=desc, - tags=d.get("tags", []), - )) + selected.append(_provider_from_discovery(d)) print(f" [+] {name:30s} {d['command']}", file=sys.stderr) if len(ollama_found) > 10: print(f" ... and {len(ollama_found) - 10} more", file=sys.stderr) @@ -514,7 +781,17 @@ def delete_provider(name: str) -> bool: return False -def call_provider(provider_name: str, prompt: str, timeout: int = 300, max_tokens: Optional[int] = None, _tried: Optional[set] = None) -> ProviderResult: +def call_provider( + provider_name: str, + prompt: str, + timeout: int = 300, + max_tokens: Optional[int] = None, + _tried: Optional[set] = None, + *, + execution_policy: Optional[ProviderExecutionPolicy] = None, + _requested_provider: Optional[str] = None, + _attempted_providers: Optional[List[str]] = None, +) -> ProviderResult: """ Call an AI provider with the given prompt. @@ -528,59 +805,110 @@ def call_provider(provider_name: str, prompt: str, timeout: int = 300, max_token timeout: Maximum execution time in seconds max_tokens: Optional max output tokens (appends provider-specific flag) _tried: Internal set of already-tried providers (prevents infinite loops) + execution_policy: Optional fallback and eligibility constraints Returns: ProviderResult with the response text or error """ + policy = execution_policy or ProviderExecutionPolicy() + requested_provider = _requested_provider or provider_name + attempted = _attempted_providers if _attempted_providers is not None else [] + if provider_name not in attempted: + attempted.append(provider_name) + observed_identity = ("mock", None, "built-in") + # Track which providers we've tried to prevent infinite fallback loops if _tried is None: _tried = set() _tried.add(provider_name) + def finish( + result: ProviderResult, + provider: Optional[Provider] = None, + ) -> ProviderResult: + """Attach facts observed by CmdForge, never values generated by the model.""" + result.requested_provider = requested_provider + result.attempted_providers = list(attempted) + if result.success: + result.actual_provider = provider_name + ( + result.model, + result.model_digest, + result.model_identity_source, + ) = observed_identity + result.locality = ( + effective_provider_locality(provider) if provider else "local" + ) + result.fallback_used = provider_name != requested_provider + return result + # Handle mock provider specially if provider_name.lower() == "mock": - return mock_provider(prompt) + return finish(mock_provider(prompt)) # Look up provider provider = get_provider(provider_name) if not provider: - return ProviderResult( + return finish(ProviderResult( text="", success=False, error=f"Provider '{provider_name}' not found. Use 'cmdforge providers' to manage providers." - ) + )) + observed_identity = runtime_provider_identity(provider) if max_tokens is not None: try: max_tokens = int(max_tokens) except (TypeError, ValueError): - return ProviderResult(text="", success=False, error="max_tokens must be an integer") + return finish(ProviderResult(text="", success=False, error="max_tokens must be an integer"), provider) if max_tokens <= 0 or max_tokens > 1_000_000: - return ProviderResult( + return finish(ProviderResult( text="", success=False, error="max_tokens must be between 1 and 1000000" - ) + ), provider) # Helper to try fallback provider(s) if available def try_fallback(error_msg: str) -> ProviderResult: import sys last_error = error_msg + if not policy.allow_fallback: + return finish( + ProviderResult(text="", success=False, error=last_error), provider + ) # Walk fallback_chain first (ordered, multi-step) if provider.fallback_chain: for fb in provider.fallback_chain: if fb not in _tried: print(f"[fallback] {provider_name} failed, trying {fb}...", file=sys.stderr) - result = call_provider(fb, prompt, timeout, max_tokens, _tried) + result = call_provider( + fb, prompt, timeout, max_tokens, _tried, + execution_policy=policy, + _requested_provider=requested_provider, + _attempted_providers=attempted, + ) if result.success: return result last_error = result.error or last_error # Fall back to single fallback (backward-compat) if provider.fallback and provider.fallback not in _tried: print(f"[fallback] {provider_name} failed, trying {provider.fallback}...", file=sys.stderr) - result = call_provider(provider.fallback, prompt, timeout, max_tokens, _tried) + result = call_provider( + provider.fallback, prompt, timeout, max_tokens, _tried, + execution_policy=policy, + _requested_provider=requested_provider, + _attempted_providers=attempted, + ) if result.success: return result last_error = result.error or last_error - return ProviderResult(text="", success=False, error=last_error) + return finish( + ProviderResult(text="", success=False, error=last_error), provider + ) + + eligibility_error = provider_policy_error( + provider, policy, observed_identity + ) + if eligibility_error: + return try_fallback(eligibility_error) # Dispatch by provider type ptype = getattr(provider, "type", None) or "subprocess" @@ -595,7 +923,7 @@ def call_provider(provider_name: str, prompt: str, timeout: int = 300, max_token return try_fallback(f"Unknown provider type: {ptype}") if result.success: - return result + return finish(result, provider) return try_fallback(result.error or "Provider call failed") except Exception as e: return try_fallback(f"Provider error: {str(e)}") diff --git a/src/cmdforge/resolver.py b/src/cmdforge/resolver.py index 901e36b..c4713b9 100644 --- a/src/cmdforge/resolver.py +++ b/src/cmdforge/resolver.py @@ -16,7 +16,7 @@ from typing import Optional, Tuple import yaml -from .tool import Tool, TOOLS_DIR, get_bin_dir, BIN_DIR +from .tool import Tool, TOOLS_DIR, get_tools_dir, get_bin_dir, BIN_DIR from .config import is_auto_fetch_enabled, load_config from .manifest import load_manifest @@ -243,7 +243,7 @@ class ToolResolver: searched_paths: list ) -> Optional[ResolvedTool]: """Search for tool in global user directory.""" - global_dir = TOOLS_DIR + global_dir = get_tools_dir() if not global_dir.exists(): return None diff --git a/src/cmdforge/runner.py b/src/cmdforge/runner.py index b792d22..1184bee 100644 --- a/src/cmdforge/runner.py +++ b/src/cmdforge/runner.py @@ -12,7 +12,9 @@ from typing import Optional, List import yaml from .tool import Tool, PromptStep, CodeStep, ToolStep, McpStep -from .providers import call_provider, mock_provider +from .providers import ( + call_provider, mock_provider, ProviderExecutionPolicy, +) from .resolver import resolve_tool, ToolNotFoundError, ToolSpec, install_from_registry from .manifest import load_manifest from .profiles import load_profile @@ -508,7 +510,9 @@ def execute_prompt_step( variables: dict, provider_override: str = None, verbose: bool = False, - base_dir: Optional[Path] = None + base_dir: Optional[Path] = None, + execution_policy: Optional[ProviderExecutionPolicy] = None, + provenance_log: Optional[list[dict]] = None, ) -> tuple[str, bool]: """ Execute a prompt step. @@ -525,6 +529,24 @@ def execute_prompt_step( Otherwise, enforces structured JSON output with schema validation. """ import re + from dataclasses import replace + + effective_policy = execution_policy or ProviderExecutionPolicy() + if step.fallback_policy == "deny" and effective_policy.allow_fallback: + effective_policy = replace(effective_policy, allow_fallback=False) + + def record_result(result, selected_provider: str) -> None: + # Mock calls bypass call_provider for deterministic tests, so attach + # the same runtime-owned identity here. + if selected_provider.lower() == "mock" and not result.actual_provider: + result.requested_provider = "mock" + result.actual_provider = "mock" + result.attempted_providers = ["mock"] + result.model = "mock" + result.locality = "local" + result.model_identity_source = "built-in" + if provenance_log is not None: + provenance_log.append(result.provenance()) # Build prompt with variable substitution try: @@ -557,7 +579,12 @@ def execute_prompt_step( if provider.lower() == "mock": result = mock_provider(prompt) else: - result = call_provider(provider, prompt, max_tokens=step.max_tokens) + result = call_provider( + provider, prompt, max_tokens=step.max_tokens, + execution_policy=effective_policy, + ) + + record_result(result, provider) if not result.success: print(f"Error in prompt step: {result.error}", file=sys.stderr) @@ -601,7 +628,12 @@ Please try again with valid JSON matching the schema exactly.""" if provider.lower() == "mock": result = mock_provider(current_prompt) else: - result = call_provider(provider, current_prompt, max_tokens=step.max_tokens) + result = call_provider( + provider, current_prompt, max_tokens=step.max_tokens, + execution_policy=effective_policy, + ) + + record_result(result, provider) if not result.success: print(f"Error in prompt step: {result.error}", file=sys.stderr) @@ -743,6 +775,8 @@ def execute_tool_step( agent_profile: Optional[str] = None, agent_skills: Optional[List[str]] = None, agent_tool_policies: Optional[List[List[str]]] = None, + execution_policy: Optional[ProviderExecutionPolicy] = None, + provenance_log: Optional[list[dict]] = None, ) -> tuple[str, bool]: """ Execute a tool step by calling another tool. @@ -815,6 +849,8 @@ def execute_tool_step( agent_tool_policies=_extend_agent_tool_policies( agent_tool_policies, step.tools ), + execution_policy=execution_policy, + provenance_log=provenance_log, ) return output, exit_code == 0 @@ -834,6 +870,8 @@ def run_tool( agent_profile: Optional[str] = None, agent_skills: Optional[List[str]] = None, agent_tool_policies: Optional[List[List[str]]] = None, + execution_policy: Optional[ProviderExecutionPolicy] = None, + provenance_log: Optional[list[dict]] = None, ) -> tuple[str, int]: """ Execute a tool. @@ -969,7 +1007,9 @@ def run_tool( variables, provider_override, verbose=verbose, - base_dir=tool_base_dir + base_dir=tool_base_dir, + execution_policy=execution_policy, + provenance_log=provenance_log, ) if not success: return "", 2 @@ -1018,6 +1058,8 @@ def run_tool( agent_profile=agent_profile, agent_skills=agent_skills, agent_tool_policies=agent_tool_policies, + execution_policy=execution_policy, + provenance_log=provenance_log, ) if not success: return "", 3 @@ -1133,14 +1175,65 @@ def create_argument_parser(tool: Tool) -> argparse.ArgumentParser: help="Show debug information") parser.add_argument("--auto-install", action="store_true", help="Automatically install missing tool dependencies") + parser.add_argument("--no-fallback", action="store_true", + help="Fail instead of trying configured fallbacks") + parser.add_argument("--require-local", action="store_true", + help="Allow only providers declared local") + parser.add_argument("--require-capability", action="append", default=[], + help="Required provider capability (repeatable)") + parser.add_argument("--data-classification", + choices=["public", "internal", "private"]) + parser.add_argument("--require-model-identity", action="store_true") + parser.add_argument("--require-model-digest", action="store_true") + parser.add_argument("--result-envelope", choices=["json"]) # Tool-specific flags from arguments for arg in tool.arguments: + existing_action = parser._option_string_actions.get(arg.flag) + if existing_action is not None: + # Preserve the legacy and useful case where a boolean tool option + # intentionally shares a universal flag (for example --verbose). + # Both consumers then read the same parsed boolean. + if arg.type == "boolean" and existing_action.dest == arg.variable: + continue + raise ValueError( + f"Tool argument {arg.flag!r} conflicts with CmdForge's " + "universal run options" + ) + + def parse_boolean(value): + if isinstance(value, bool): + return value + normalized = str(value).strip().lower() + if normalized in {"1", "true", "yes", "on"}: + return True + if normalized in {"0", "false", "no", "off"}: + return False + raise argparse.ArgumentTypeError( + "expected true/false, yes/no, on/off, or 1/0" + ) + + value_type = { + "string": str, + "integer": int, + "number": float, + "boolean": parse_boolean, + }[arg.type] + options = { + "dest": arg.variable, + "default": arg.default, + "help": arg.description or f"{arg.variable} (default: {arg.default})", + "required": arg.required, + "type": value_type, + } + if arg.enum is not None: + options["choices"] = [value_type(value) for value in arg.enum] + if arg.type == "boolean": + options["nargs"] = "?" + options["const"] = True parser.add_argument( arg.flag, - dest=arg.variable, - default=arg.default, - help=arg.description or f"{arg.variable} (default: {arg.default})" + **options, ) return parser @@ -1215,6 +1308,16 @@ def main(): # Determine provider override (CLI flag takes precedence over manifest) effective_provider = args.provider or provider_override_from_manifest + execution_policy = ProviderExecutionPolicy( + allow_fallback=not args.no_fallback, + required_locality="local" if args.require_local else None, + required_capabilities=tuple(args.require_capability or []), + data_classification=args.data_classification or "public", + require_model_identity=args.require_model_identity, + require_model_digest=args.require_model_digest, + ) + provider_calls = [] + # Run tool output, exit_code = run_tool( tool=tool, @@ -1224,15 +1327,36 @@ def main(): dry_run=args.dry_run, show_prompt=args.show_prompt, verbose=args.verbose, - auto_install=args.auto_install + auto_install=args.auto_install, + execution_policy=execution_policy, + provenance_log=provider_calls, ) # Write output - if exit_code == 0 and output: - if args.output_file: - Path(args.output_file).write_text(output) + if args.result_envelope == "json" or (exit_code == 0 and output): + if args.result_envelope == "json": + import json + rendered_output = json.dumps({ + "output": output, + "execution": { + "exit_code": exit_code, + "provider_calls": provider_calls, + "fallback_used": any( + item.get("fallback_used", False) + for item in provider_calls + ), + "actual_providers": list(dict.fromkeys( + item["actual_provider"] for item in provider_calls + if item.get("actual_provider") + )), + }, + }, indent=2) else: - print(output) + rendered_output = output + if args.output_file: + Path(args.output_file).write_text(rendered_output) + else: + print(rendered_output) if exit_code == 0 and not args.dry_run: try: diff --git a/src/cmdforge/tool.py b/src/cmdforge/tool.py index 97148af..b46588a 100644 --- a/src/cmdforge/tool.py +++ b/src/cmdforge/tool.py @@ -124,6 +124,15 @@ class ToolArgument: required: bool = False def __post_init__(self) -> None: + # Older CmdForge tools used Python-style scalar names. Normalize them + # on load so typed parsing does not break existing user tools. + legacy_type_aliases = { + "str": "string", + "int": "integer", + "float": "number", + "bool": "boolean", + } + self.type = legacy_type_aliases.get(self.type, self.type) if self.type not in ("string", "integer", "number", "boolean"): raise ValueError( "ToolArgument type must be string, integer, number, or boolean" @@ -178,9 +187,12 @@ class PromptStep: plain_text: bool = False # Bypass structured output enforcement max_tokens: Optional[int] = None # Max output tokens (provider-dependent) skills: Optional[List[str]] = None # Skill names to enable for this step + fallback_policy: str = "allow" # "allow" | "deny" def __post_init__(self) -> None: _validate_skill_selection(self.skills) + if self.fallback_policy not in {"allow", "deny"}: + raise ValueError("PromptStep fallback_policy must be allow or deny") def to_dict(self) -> dict: d = { @@ -207,6 +219,8 @@ class PromptStep: d["max_tokens"] = self.max_tokens if self.skills is not None: d["skills"] = self.skills + if self.fallback_policy != "allow": + d["fallback_policy"] = self.fallback_policy return d @classmethod @@ -233,6 +247,7 @@ class PromptStep: plain_text=data.get("plain_text", False), max_tokens=max_tokens, skills=data.get("skills"), + fallback_policy=data.get("fallback_policy", "allow"), ) @@ -598,6 +613,11 @@ def get_tools_dir() -> Path: return TOOLS_DIR +def get_project_tools_dir(project_dir: Optional[Path] = None) -> Path: + """Return the canonical project-local tool root without creating it.""" + return (project_dir or Path.cwd()).resolve() / ".cmdforge" + + def ensure_settings(tool_dir: Path) -> Optional[Path]: """Ensure settings.yaml exists if defaults.yaml exists. @@ -747,12 +767,32 @@ def load_tool(name: str) -> Optional[Tool]: _REGISTRY_FIELDS = ("registry_hash", "registry_status", "registry_owner", "registry_feedback") -def save_tool(tool: Tool) -> Path: - """Save a tool to disk, preserving registry metadata from existing config.""" - tool_dir = get_tools_dir() / tool.name +def save_tool( + tool: Tool, + *, + tools_dir: Optional[Path] = None, + create_wrapper: Optional[bool] = None, +) -> Path: + """Save a tool globally or under an explicit project-local tool root. + + Global saves keep the historical wrapper behavior. Project-local saves are + invoked through ``cmdforge run`` so a global wrapper cannot accidentally + resolve a different project tool with the same name. + """ + valid, error = validate_tool_name(tool.name) + if not valid: + raise ValueError(error) + global_tools_dir = get_tools_dir() + target_root = Path(tools_dir).resolve() if tools_dir is not None else global_tools_dir + target_root.mkdir(parents=True, exist_ok=True) + tool_dir = target_root / tool.name + if tool_dir.is_symlink(): + raise ValueError(f"Refusing to save through symlinked tool directory: {tool_dir}") tool_dir.mkdir(parents=True, exist_ok=True) config_path = tool_dir / "config.yaml" + if config_path.is_symlink(): + raise ValueError(f"Refusing to overwrite symlinked tool config: {config_path}") # Preserve registry fields from existing config (not part of Tool model) preserved = {} @@ -769,8 +809,12 @@ def save_tool(tool: Tool) -> Path: new_data.update(preserved) config_path.write_text(yaml.dump(new_data, default_flow_style=False, sort_keys=False)) - # Create wrapper script - create_wrapper_script(tool.name) + # Global tools retain convenient shell wrappers. Project tools are + # intentionally cwd-scoped and must be resolved by ``cmdforge run``. + if create_wrapper is None: + create_wrapper = target_root == global_tools_dir.resolve() + if create_wrapper: + create_wrapper_script(tool.name) # Ensure settings.yaml exists if defaults.yaml exists ensure_settings(tool_dir) diff --git a/tests/test_cli.py b/tests/test_cli.py index 9468af4..a95d2de 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -157,6 +157,33 @@ class TestListCommand: assert main() == 1 assert "--full requires --json" in capsys.readouterr().err + def test_list_includes_project_tools_and_prefers_project_shadow( + self, temp_tools_dir, tmp_path, monkeypatch, capsys + ): + from cmdforge.tool import save_tool + + save_tool(Tool(name="shadowed", description="global")) + project = tmp_path / "project" + project.mkdir() + save_tool( + Tool(name="shadowed", description="project"), + tools_dir=project / ".cmdforge", + ) + save_tool( + Tool(name="project-only", description="local"), + tools_dir=project / ".cmdforge", + ) + monkeypatch.chdir(project) + + with patch("sys.argv", ["cmdforge", "list", "--json"]): + assert main() == 0 + + payload = json.loads(capsys.readouterr().out) + by_name = {item["name"]: item for item in payload} + assert by_name["shadowed"]["description"] == "project" + assert "project-only" in by_name + assert [item["name"] for item in payload].count("shadowed") == 1 + class TestRunOnceCommand: def test_combines_instruction_with_piped_input(self, capsys): @@ -243,6 +270,29 @@ class TestRunOnceCommand: call_provider.assert_not_called() assert message in capsys.readouterr().err + def test_json_envelope_contains_runtime_provenance(self, capsys): + from cmdforge.providers import ProviderResult + + result = ProviderResult( + text="answer", success=True, + requested_provider="local", actual_provider="local", + attempted_providers=["local"], locality="local", + ) + with ( + patch("sys.argv", [ + "cmdforge", "run-once", "hello", "--provider", "local", + "--no-fallback", "--result-envelope", "json", + ]), + patch("sys.stdin", StringIO()), + patch("cmdforge.providers.call_provider", return_value=result), + ): + assert main() == 0 + + payload = json.loads(capsys.readouterr().out) + call = payload["execution"]["provider_calls"][0] + assert call["actual_provider"] == "local" + assert call["fallback_used"] is False + class TestCreateCommand: """Tests for 'cmdforge create' command.""" @@ -296,6 +346,36 @@ class TestCreateCommand: captured = capsys.readouterr() assert 'invalid' in captured.out.lower() or 'invalid' in captured.err.lower() + def test_create_project_tool_does_not_create_global_wrapper( + self, temp_tools_dir, tmp_path, monkeypatch, capsys + ): + project = tmp_path / "project" + project.mkdir() + monkeypatch.chdir(project) + + with patch("sys.argv", [ + "cmdforge", "create", "project-echo", "--project", + "--prompt", "Echo {input}", "--provider", "mock", + ]): + assert main() == 0 + + config = project / ".cmdforge" / "project-echo" / "config.yaml" + assert config.exists() + assert not (tmp_path / ".cmdforge" / "project-echo" / "config.yaml").exists() + assert not (tmp_path / ".local" / "bin" / "project-echo").exists() + assert "cmdforge run project-echo" in capsys.readouterr().out + + def test_create_output_dir_uses_explicit_tool_root( + self, temp_tools_dir, tmp_path + ): + target = tmp_path / "portable-tools" + with patch("sys.argv", [ + "cmdforge", "create", "portable", "--output-dir", str(target), + ]): + assert main() == 0 + + assert (target / "portable" / "config.yaml").exists() + class TestDeleteCommand: """Tests for 'cmdforge delete' command.""" @@ -405,6 +485,91 @@ class TestRunCommand: assert result != 0 + def test_run_prefers_project_local_tool( + self, temp_tools_dir, tmp_path, monkeypatch, capsys + ): + from cmdforge.tool import save_tool + + save_tool(Tool(name="shadowed", output="global")) + project = tmp_path / "project" + project.mkdir() + save_tool( + Tool(name="shadowed", output="project"), + tools_dir=project / ".cmdforge", + ) + monkeypatch.chdir(project) + + with ( + patch("sys.argv", ["cmdforge", "run", "shadowed"]), + patch("sys.stdin", StringIO("")), + patch("sys.stdin.isatty", return_value=True), + ): + assert main() == 0 + + assert capsys.readouterr().out.strip() == "project" + + def test_run_json_envelope_preserves_output_and_provenance( + self, temp_tools_dir, capsys + ): + from cmdforge.tool import save_tool + + save_tool(Tool( + name="audited", + steps=[PromptStep( + prompt="Echo {input}", provider="mock", + output_var="answer", plain_text=True, + )], + output="{answer}", + )) + + with ( + patch("sys.argv", [ + "cmdforge", "run", "audited", "--result-envelope", "json", + ]), + patch("sys.stdin", StringIO("hello")), + patch("sys.stdin.isatty", return_value=False), + ): + assert main() == 0 + + payload = json.loads(capsys.readouterr().out) + assert "MOCK" in payload["output"] + assert payload["execution"]["actual_providers"] == ["mock"] + assert payload["execution"]["provider_calls"][0]["locality"] == "local" + + def test_failed_run_still_emits_audit_envelope( + self, temp_tools_dir, capsys + ): + from cmdforge.providers import ProviderResult + from cmdforge.tool import save_tool + + save_tool(Tool( + name="strict", + steps=[PromptStep( + prompt="Private", provider="primary", + output_var="answer", plain_text=True, + )], + output="{answer}", + )) + failure = ProviderResult( + text="", success=False, error="offline", + requested_provider="primary", attempted_providers=["primary"], + ) + with ( + patch("sys.argv", [ + "cmdforge", "run", "strict", "--no-fallback", + "--result-envelope", "json", + ]), + patch("sys.stdin", StringIO("private")), + patch("sys.stdin.isatty", return_value=False), + patch("cmdforge.runner.call_provider", return_value=failure), + ): + assert main() == 2 + + payload = json.loads(capsys.readouterr().out) + assert payload["output"] == "" + assert payload["execution"]["exit_code"] == 2 + assert payload["execution"]["provider_calls"][0]["error"] == "offline" + class TestTestCommand: """Tests for 'cmdforge test' command.""" @@ -458,6 +623,29 @@ class TestProvidersCommand: # Should show some default providers assert 'mock' in captured.out.lower() or 'claude' in captured.out.lower() + def test_providers_list_json_exposes_routing_metadata( + self, temp_providers_file, capsys + ): + from cmdforge.providers import Provider, save_providers + + save_providers([Provider( + "safe", "safe-cli", locality="local", + capabilities=["structured-json"], data_policy="private", + max_context_tokens=8192, + )]) + with ( + patch("sys.argv", ["cmdforge", "providers", "list", "--json"]), + patch("cmdforge.cli.provider_commands._provider_status", return_value=(True, "Ready")), + ): + assert main() == 0 + + payload = json.loads(capsys.readouterr().out) + assert payload[0]["locality"] == "local" + assert payload[0]["capabilities"] == ["structured-json"] + assert payload[0]["data_policy"] == "private" + assert payload[0]["max_context_tokens"] == 8192 + assert payload[0]["available"] is True + def test_providers_add(self, temp_providers_file, capsys): """Add a custom provider.""" with patch('sys.argv', ['cmdforge', 'providers', 'add', @@ -479,6 +667,13 @@ class TestProvidersCommand: '--model', 'example/model', '--api-key-env', 'CUSTOM_API_KEY', '--tag', 'api', + '--locality', 'remote', + '--capability', 'structured-json', + '--model-digest', 'sha256:abc', + '--cost-class', 'low', + '--latency-class', 'fast', + '--max-context-tokens', '128000', + '--data-policy', 'internal', '--fallback-chain', 'free', ]): result = main() @@ -491,6 +686,13 @@ class TestProvidersCommand: assert provider.api_key_env == 'CUSTOM_API_KEY' assert provider.tags == ['api'] assert provider.fallback_chain == PRESET_CHAINS['free'] + assert provider.locality == 'remote' + assert provider.capabilities == ['structured-json'] + assert provider.model_digest == 'sha256:abc' + assert provider.cost_class == 'low' + assert provider.latency_class == 'fast' + assert provider.max_context_tokens == 128000 + assert provider.data_policy == 'internal' def test_providers_list_reports_missing_api_key(self, temp_providers_file, capsys): from cmdforge.providers import Provider, save_providers @@ -828,6 +1030,23 @@ def test_switch_to_existing_tool_closes_creation_page_first(): ] +def test_tool_builder_save_scope_prefers_recognized_projects(tmp_path): + pytest.importorskip("PySide6") + from cmdforge.gui.pages.tool_builder_page import default_save_scope + + assert default_save_scope(tmp_path) == "global" + (tmp_path / "pyproject.toml").write_text("[project]\nname='demo'\n") + assert default_save_scope(tmp_path) == "project" + + +def test_tool_builder_existing_tool_root_does_not_nest_tool_name(tmp_path): + pytest.importorskip("PySide6") + from cmdforge.gui.pages.tool_builder_page import existing_tool_root + + config = tmp_path / ".cmdforge" / "demo" / "config.yaml" + assert existing_tool_root(config) == (tmp_path / ".cmdforge").resolve() + + def test_guided_extraction_builds_new_tool_without_mutating_draft(): pytest.importorskip("PySide6") from cmdforge.gui.pages.tool_builder_page import ( diff --git a/tests/test_providers.py b/tests/test_providers.py index 071cfd7..84d3215 100644 --- a/tests/test_providers.py +++ b/tests/test_providers.py @@ -8,7 +8,7 @@ import pytest import yaml from cmdforge.providers import ( - Provider, ProviderResult, + Provider, ProviderResult, ProviderExecutionPolicy, load_providers, save_providers, get_provider, add_provider, delete_provider, call_provider, discover_installed_providers, mock_provider, @@ -60,6 +60,19 @@ class TestProvider: provider = Provider.from_dict(data) assert provider.description == "" + def test_legacy_policy_field_aliases_are_canonicalized(self): + provider = Provider.from_dict({ + "name": "legacy-policy", + "command": "provider", + "model_id": "model-a", + "max_input_tokens": 8192, + }) + + assert provider.model == "model-a" + assert provider.max_context_tokens == 8192 + assert provider.to_dict()["model"] == "model-a" + assert provider.to_dict()["max_context_tokens"] == 8192 + def test_roundtrip(self): original = Provider( name="custom", @@ -90,6 +103,13 @@ class TestProvider: fallback_chain=["backup", "mock"], api_key_env="TEST_API_KEY", pty_config={"prompt_pattern": ">"}, + locality="remote", + capabilities=["text", "structured-json"], + model_digest="sha256:abc", + cost_class="low", + latency_class="fast", + max_context_tokens=131072, + data_policy="internal", ) assert Provider.from_dict(original.to_dict()) == original @@ -516,6 +536,200 @@ class TestProviderFallback: # Fallback to mock should succeed assert result.success is True assert "[MOCK]" in result.text + assert result.requested_provider == "primary" + assert result.actual_provider == "mock" + assert result.attempted_providers == ["primary", "mock"] + assert result.fallback_used is True + + def test_no_fallback_fails_closed(self, temp_providers_file): + save_providers([ + Provider("primary", "nonexistent-cmd", fallback="mock"), + Provider("mock", "mock"), + ]) + + result = call_provider( + "primary", "Test prompt", + execution_policy=ProviderExecutionPolicy(allow_fallback=False), + ) + + assert result.success is False + assert result.requested_provider == "primary" + assert result.actual_provider is None + assert result.attempted_providers == ["primary"] + assert result.fallback_used is False + + def test_private_data_requires_explicit_compatible_policy( + self, temp_providers_file + ): + save_providers([ + Provider( + "local", "unused", locality="local", + capabilities=["structured-json"], + ), + ]) + + result = call_provider( + "local", "private packet", + execution_policy=ProviderExecutionPolicy( + allow_fallback=False, + required_locality="local", + required_capabilities=("structured-json",), + data_classification="private", + ), + ) + + assert result.success is False + assert "data_policy is unspecified" in result.error + + def test_require_local_rejects_ollama_redirected_to_remote_host( + self, temp_providers_file, monkeypatch + ): + save_providers([ + Provider( + "ollama-private", "ollama run safe", locality="local", + data_policy="private", + ), + ]) + monkeypatch.setenv("OLLAMA_HOST", "gpu.example.test:11434") + + result = call_provider( + "ollama-private", "private packet", + execution_policy=ProviderExecutionPolicy( + allow_fallback=False, + required_locality="local", + data_classification="private", + ), + ) + + assert result.success is False + assert "locality is remote" in result.error + + def test_ollama_digest_is_observed_from_local_api( + self, temp_providers_file, monkeypatch + ): + save_providers([ + Provider( + "ollama-safe", "ollama run safe:latest", model="safe:latest", + locality="local", data_policy="private", + ), + ]) + monkeypatch.delenv("OLLAMA_HOST", raising=False) + + class Response: + def __enter__(self): + return self + + def __exit__(self, *args): + return False + + def read(self): + return ( + b'{"models":[{"name":"safe:latest","model":"safe:latest",' + b'"digest":"abcdef"}]}' + ) + + with ( + patch("urllib.request.urlopen", return_value=Response()), + patch( + "cmdforge.providers.call_provider_subprocess", + return_value=ProviderResult(text="safe", success=True), + ), + ): + result = call_provider( + "ollama-safe", "private", + execution_policy=ProviderExecutionPolicy( + allow_fallback=False, + required_locality="local", + data_classification="private", + require_model_identity=True, + require_model_digest=True, + ), + ) + + assert result.model == "safe:latest" + assert result.model_digest == "sha256:abcdef" + assert result.model_identity_source == "ollama-api" + + def test_policy_eligible_provider_reports_verified_metadata( + self, temp_providers_file + ): + provider = Provider( + "local", "local-ai", model="model-a", + model_digest="sha256:123", locality="local", + capabilities=["structured-json"], data_policy="private", + ) + save_providers([provider]) + low_level = ProviderResult(text='{"ok": true}', success=True) + + with patch( + "cmdforge.providers.call_provider_subprocess", + return_value=low_level, + ): + result = call_provider( + "local", "private packet", + execution_policy=ProviderExecutionPolicy( + allow_fallback=False, + required_locality="local", + required_capabilities=("structured-json",), + data_classification="private", + require_model_identity=True, + require_model_digest=True, + ), + ) + + assert result.success is True + assert result.provenance() == { + "success": True, + "requested_provider": "local", + "actual_provider": "local", + "attempted_providers": ["local"], + "fallback_used": False, + "model": "model-a", + "model_digest": "sha256:123", + "locality": "local", + "model_identity_source": "provider-config", + } + + def test_fallback_chain_skips_policy_ineligible_provider( + self, temp_providers_file + ): + save_providers([ + Provider( + "primary", "primary", fallback_chain=["remote", "local"], + locality="local", data_policy="private", + ), + Provider( + "remote", "remote", locality="remote", data_policy="private", + ), + Provider( + "local", "local", locality="local", data_policy="private", + model="safe-model", + ), + ]) + + def invoke(provider, prompt, timeout, max_tokens): + if provider.name == "primary": + return ProviderResult(text="", success=False, error="offline") + assert provider.name == "local" + return ProviderResult(text="safe", success=True) + + with patch( + "cmdforge.providers.call_provider_subprocess", side_effect=invoke + ) as low_level: + result = call_provider( + "primary", "private packet", + execution_policy=ProviderExecutionPolicy( + required_locality="local", + data_classification="private", + ), + ) + + assert result.success is True + assert result.actual_provider == "local" + assert result.attempted_providers == ["primary", "remote", "local"] + assert [call.args[0].name for call in low_level.call_args_list] == [ + "primary", "local" + ] def test_fallback_chain(self, temp_providers_file): """Fallback can chain to another provider with fallback.""" diff --git a/tests/test_runner.py b/tests/test_runner.py index 0bfdad8..55b1887 100644 --- a/tests/test_runner.py +++ b/tests/test_runner.py @@ -229,6 +229,23 @@ class TestExecutePromptStep: assert success is False assert output == "" + @patch('cmdforge.runner.call_provider') + def test_step_can_deny_fallback_independently_of_caller(self, mock_call): + mock_call.return_value = ProviderResult( + text="private", success=True, + ) + step = PromptStep( + prompt="Private", provider="local", output_var="out", + plain_text=True, fallback_policy="deny", + ) + + output, success = execute_prompt_step(step, {"input": ""}) + + assert success is True + assert output == "private" + policy = mock_call.call_args.kwargs["execution_policy"] + assert policy.allow_fallback is False + @patch('cmdforge.runner.mock_provider') def test_mock_provider_used(self, mock_mock): mock_mock.return_value = ProviderResult(text="mock response", success=True) @@ -988,6 +1005,58 @@ class TestCreateArgumentParser: assert args.count == "10" + def test_parser_enforces_typed_boolean_integer_and_enum_arguments(self): + tool = Tool( + name="typed", + arguments=[ + ToolArgument( + flag="--project", variable="project", + type="boolean", default=False, + ), + ToolArgument( + flag="--count", variable="count", + type="integer", required=True, + ), + ToolArgument( + flag="--mode", variable="mode", + enum=["safe", "fast"], default="safe", + ), + ], + ) + parser = create_argument_parser(tool) + + args = parser.parse_args(["--project", "--count", "3", "--mode", "fast"]) + + assert args.project is True + assert args.count == 3 + assert args.mode == "fast" + + def test_parser_reuses_matching_legacy_boolean_universal_flag(self): + tool = Tool( + name="legacy", + arguments=[ + ToolArgument( + flag="--verbose", variable="verbose", type="bool" + ) + ], + ) + + args = create_argument_parser(tool).parse_args(["--verbose"]) + + assert args.verbose is True + assert collect_custom_args(tool, args) == {"verbose": True} + + def test_parser_rejects_incompatible_universal_flag_collision(self): + tool = Tool( + name="conflict", + arguments=[ + ToolArgument(flag="--provider", variable="destination") + ], + ) + + with pytest.raises(ValueError, match="universal run options"): + create_argument_parser(tool) + def test_parser_input_output_flags(self): tool = Tool(name="test") parser = create_argument_parser(tool) diff --git a/tests/test_skills.py b/tests/test_skills.py index 3d975ef..6bc1408 100644 --- a/tests/test_skills.py +++ b/tests/test_skills.py @@ -378,7 +378,7 @@ class TestRunnerSkillIntegration: ): captured = {} - def fake_call(provider, prompt, max_tokens=None): + def fake_call(provider, prompt, max_tokens=None, **kwargs): captured.update(provider=provider, prompt=prompt) return ProviderResult(text="done", success=True) @@ -412,7 +412,7 @@ class TestRunnerSkillIntegration: ): captured = {} - def fake_call(provider, prompt, max_tokens=None): + def fake_call(provider, prompt, max_tokens=None, **kwargs): captured.update(provider=provider, prompt=prompt) return ProviderResult(text='{"output": "done"}', success=True) diff --git a/tests/test_tool.py b/tests/test_tool.py index b2282b7..c625b3e 100644 --- a/tests/test_tool.py +++ b/tests/test_tool.py @@ -105,6 +105,25 @@ class TestToolArgument: with pytest.raises(ValueError): ToolArgument(flag="--value", variable="value", **values) + @pytest.mark.parametrize( + ("legacy_type", "canonical_type"), + [ + ("str", "string"), + ("int", "integer"), + ("float", "number"), + ("bool", "boolean"), + ], + ) + def test_legacy_scalar_type_aliases_are_normalized( + self, legacy_type, canonical_type + ): + argument = ToolArgument( + flag="--value", variable="value", type=legacy_type + ) + + assert argument.type == canonical_type + assert argument.to_dict().get("type", "string") == canonical_type + class TestPromptStep: """Tests for PromptStep dataclass.""" @@ -686,6 +705,23 @@ class TestAgentContext: restored = PromptStep.from_dict(step.to_dict()) assert restored.skills == ["python"] + def test_promptstep_deny_fallback_roundtrip(self): + step = PromptStep( + prompt="Private", provider="local", output_var="out", + fallback_policy="deny", + ) + + restored = PromptStep.from_dict(step.to_dict()) + + assert restored.fallback_policy == "deny" + + def test_promptstep_rejects_invalid_fallback_policy(self): + with pytest.raises(ValueError, match="fallback_policy"): + PromptStep( + prompt="Test", provider="local", output_var="out", + fallback_policy="sometimes", + ) + class TestDeprecationFields: def test_round_trip_and_warning(self):