Add provider fallback, visual schema builder, and UI improvements

Provider fallback system:
- Add optional fallback field to Provider dataclass
- Automatically try fallback provider when primary fails (timeout, error, offline)
- Prevent infinite loops with circular fallback detection
- Add fallback column and dropdown in providers UI

Visual schema builder for prompt steps:
- Add SchemaBuilderDialog with table-based field editor
- Support field name, type, description, and required flag
- Live JSON schema preview
- Integrate into PromptStepDialog with Edit Schema button
- Fix output_schema being lost when editing prompt steps

Other improvements:
- Add search box to tools page
- Improve schema example generation for structured output
- Update CLAUDE.md documentation

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>
This commit is contained in:
rob 2026-02-21 00:50:18 -04:00
parent 8581775002
commit 71c6358b9d
9 changed files with 892 additions and 99 deletions

View File

@ -15,14 +15,12 @@ pip install -e ".[dev]"
# Run all unit tests (excluding integration tests that need a server)
pytest tests/ -m "not integration"
# Run all tests with verbose output
pytest tests/ -v
# Run a specific test file
pytest tests/test_runner.py -v
# Run a specific test class
# Run a specific test class or method
pytest tests/test_runner.py::TestSubstituteVariables -v
pytest tests/test_runner.py::TestSubstituteVariables::test_simple_substitution -v
# Run with coverage
pytest tests/ --cov=cmdforge --cov-report=html
@ -31,11 +29,10 @@ pytest tests/ --cov=cmdforge --cov-report=html
python -m cmdforge.registry.app # Start server first
pytest tests/test_registry_integration.py -v -m integration
# Run the CLI
python -m cmdforge.cli
# Launch the GUI
cmdforge
# CLI entry points
cmdforge # Main CLI / GUI launcher
cf # Interactive tool picker (fzf-style)
python -m cmdforge.cli # Alternative CLI invocation
```
## Architecture
@ -76,7 +73,10 @@ Tools are YAML configs with:
### Step Types
1. **Prompt Step**: Calls AI provider with template, stores result in `output_var`. Supports `profile` for AI personas and `strip_fences` for markdown cleanup
1. **Prompt Step**: Calls AI provider with template, stores result in `output_var`
- `profile`: AI persona (system prompt prefix)
- `strip_fences`: Remove markdown code fences from output
- `structured_output`: JSON schema for validated structured responses
2. **Code Step**: Executes Python code via `exec()`, captures specified variables (comma-separated for multiple outputs)
3. **Tool Step**: Calls another tool (meta-tools), supports `args` dict and `provider` override. Dependencies resolved via `resolver.py`
@ -103,7 +103,7 @@ providers:
command: "echo '[MOCK]'"
```
The `mock` provider is built-in for testing without API calls.
The `mock` provider is built-in for testing without API calls. Use `--provider mock` or `--dry-run` flags when testing tools.
## Web UI & Registry
@ -145,14 +145,22 @@ CMDFORGE_REGISTRY_DB=/path/to/db PORT=5050 python -m cmdforge.web.app
## Testing Conventions
Tests use `pytest` with common fixtures:
- `temp_tools_dir(tmp_path)`: Redirects `TOOLS_DIR` and `BIN_DIR` to temp directory
- `temp_providers_file(tmp_path)`: Redirects `providers.yaml` to temp directory
Tests use `pytest` without a shared conftest.py. Common patterns:
Mocking strategy:
- **File system**: Use `tmp_path` fixture and patch module-level paths
**Mocking strategy:**
- **File system**: Use `tmp_path` fixture and patch module-level paths (`TOOLS_DIR`, `BIN_DIR`, `PROVIDERS_FILE`)
- **Subprocess calls**: Mock `subprocess.run` and `shutil.which`
- **Provider calls**: Mock `call_provider` to return `ProviderResult`
- **Provider calls**: Mock `call_provider` to return `ProviderResult(text="...", success=True)`
- **Registry calls**: Mock `requests` for API tests
**Example fixture pattern:**
```python
@pytest.fixture
def temp_providers_file(tmp_path):
providers_file = tmp_path / "providers.yaml"
with patch('cmdforge.providers.PROVIDERS_FILE', providers_file):
yield providers_file
```
Integration tests are marked with `@pytest.mark.integration` and require a running registry server.

View File

@ -2,10 +2,10 @@
from PySide6.QtWidgets import (
QDialog, QVBoxLayout, QFormLayout, QLineEdit,
QPushButton, QHBoxLayout, QLabel
QPushButton, QHBoxLayout, QLabel, QComboBox
)
from ...providers import Provider, add_provider
from ...providers import Provider, add_provider, load_providers
class ProviderDialog(QDialog):
@ -43,12 +43,20 @@ class ProviderDialog(QDialog):
self.desc_input.setPlaceholderText("Claude AI via claude-cli")
form.addRow("Description:", self.desc_input)
self.fallback_combo = QComboBox()
self.fallback_combo.addItem("(none)", None)
# Populate with available providers
for p in sorted(load_providers(), key=lambda x: x.name):
self.fallback_combo.addItem(p.name, p.name)
form.addRow("Fallback:", self.fallback_combo)
layout.addLayout(form)
# Help text
help_text = QLabel(
"The command should accept input on stdin and output to stdout.\n"
"Example commands: claude-cli, sgpt, llm, mods"
"Example commands: claude-cli, sgpt, llm, mods\n"
"Fallback: Provider to use if this one fails (timeout, error, offline)"
)
help_text.setStyleSheet("color: #718096; font-size: 11px;")
help_text.setWordWrap(True)
@ -77,6 +85,11 @@ class ProviderDialog(QDialog):
self.name_input.setEnabled(False) # Can't rename
self.cmd_input.setText(provider.command)
self.desc_input.setText(provider.description or "")
# Set fallback selection
if provider.fallback:
idx = self.fallback_combo.findData(provider.fallback)
if idx >= 0:
self.fallback_combo.setCurrentIndex(idx)
def _save(self):
"""Save the provider."""
@ -91,10 +104,17 @@ class ProviderDialog(QDialog):
self.cmd_input.setFocus()
return
description = self.desc_input.text().strip() or None
description = self.desc_input.text().strip()
fallback = self.fallback_combo.currentData()
# Prevent self-referential fallback
if fallback == name:
from PySide6.QtWidgets import QMessageBox
QMessageBox.warning(self, "Invalid", "A provider cannot fall back to itself.")
return
try:
add_provider(name, command, description)
add_provider(Provider(name, command, description, fallback))
self.accept()
except Exception as e:
from PySide6.QtWidgets import QMessageBox

View File

@ -1,12 +1,14 @@
"""Step editor dialogs."""
import ast
import json
from PySide6.QtWidgets import (
QDialog, QVBoxLayout, QFormLayout, QLineEdit,
QComboBox, QPushButton, QHBoxLayout, QLabel,
QPlainTextEdit, QSplitter, QGroupBox, QTextEdit, QMessageBox,
QCheckBox, QSpinBox
QCheckBox, QSpinBox, QTableWidget, QTableWidgetItem, QHeaderView,
QWidget, QAbstractItemView
)
from PySide6.QtCore import Qt, QThread, Signal
@ -15,14 +17,260 @@ from ...providers import load_providers, call_provider
from ...profiles import list_profiles
class SchemaBuilderDialog(QDialog):
"""Visual builder for JSON output schemas."""
TYPES = ["string", "number", "integer", "boolean", "array", "object"]
def __init__(self, parent, schema: dict = None):
super().__init__(parent)
self.setWindowTitle("Output Schema Builder")
self.setMinimumSize(850, 500)
self._setup_ui()
if schema:
self._load_schema(schema)
def _setup_ui(self):
"""Set up the UI."""
layout = QVBoxLayout(self)
layout.setSpacing(12)
# Instructions
info = QLabel(
"Define the fields that the AI should return. "
"The AI will be instructed to respond with JSON matching this schema."
)
info.setWordWrap(True)
info.setStyleSheet("color: #718096;")
layout.addWidget(info)
# Fields table
self.table = QTableWidget()
self.table.setColumnCount(5)
self.table.setHorizontalHeaderLabels(["Field Name", "Type", "Description", "Required", "Delete"])
self.table.horizontalHeader().setSectionResizeMode(0, QHeaderView.Interactive)
self.table.horizontalHeader().setSectionResizeMode(1, QHeaderView.Interactive)
self.table.horizontalHeader().setSectionResizeMode(2, QHeaderView.Stretch)
self.table.horizontalHeader().setSectionResizeMode(3, QHeaderView.Fixed)
self.table.horizontalHeader().setSectionResizeMode(4, QHeaderView.Fixed)
self.table.setColumnWidth(0, 150) # Field Name
self.table.setColumnWidth(1, 120) # Type
self.table.setColumnWidth(3, 80) # Required checkbox
self.table.setColumnWidth(4, 100) # Delete button
self.table.verticalHeader().setDefaultSectionSize(50) # Row height
self.table.setSelectionMode(QAbstractItemView.NoSelection) # No selection needed
self.table.setFocusPolicy(Qt.NoFocus) # Remove focus rectangle
self.table.verticalHeader().setVisible(False)
layout.addWidget(self.table, 1)
# Add field button
btn_row = QHBoxLayout()
self.btn_add = QPushButton("+ Add Field")
self.btn_add.clicked.connect(self._add_field)
btn_row.addWidget(self.btn_add)
btn_row.addStretch()
layout.addLayout(btn_row)
# Preview section
preview_group = QGroupBox("Schema Preview (JSON)")
preview_layout = QVBoxLayout(preview_group)
self.preview = QPlainTextEdit()
self.preview.setReadOnly(True)
self.preview.setMaximumHeight(120)
font = self.preview.font()
font.setFamily("Consolas, Monaco, monospace")
font.setPointSize(9)
self.preview.setFont(font)
preview_layout.addWidget(self.preview)
layout.addWidget(preview_group)
# Buttons
buttons = QHBoxLayout()
self.btn_clear = QPushButton("Clear All")
self.btn_clear.clicked.connect(self._clear_all)
buttons.addWidget(self.btn_clear)
buttons.addStretch()
self.btn_cancel = QPushButton("Cancel")
self.btn_cancel.setObjectName("secondary")
self.btn_cancel.clicked.connect(self.reject)
buttons.addWidget(self.btn_cancel)
self.btn_ok = QPushButton("Apply Schema")
self.btn_ok.clicked.connect(self.accept)
buttons.addWidget(self.btn_ok)
layout.addLayout(buttons)
# Add default fields for new schema
if not self.table.rowCount():
self._add_field("output", "string", "The main output content", True)
self._add_field("reasoning", "string", "Explanation of the response", False)
self._update_preview()
def _add_field(self, name: str = "", field_type: str = "string",
description: str = "", required: bool = True):
"""Add a field row to the table."""
row = self.table.rowCount()
self.table.insertRow(row)
# Field name
name_input = QLineEdit(name)
name_input.setPlaceholderText("field_name")
name_input.setMinimumHeight(36)
name_input.textChanged.connect(self._update_preview)
self.table.setCellWidget(row, 0, name_input)
# Type combo
type_combo = QComboBox()
type_combo.addItems(self.TYPES)
type_combo.setMinimumHeight(36)
if field_type in self.TYPES:
type_combo.setCurrentText(field_type)
type_combo.currentTextChanged.connect(self._update_preview)
self.table.setCellWidget(row, 1, type_combo)
# Description
desc_input = QLineEdit(description)
desc_input.setPlaceholderText("Description of this field")
desc_input.setMinimumHeight(36)
desc_input.textChanged.connect(self._update_preview)
self.table.setCellWidget(row, 2, desc_input)
# Required checkbox - center it
req_widget = QWidget()
req_layout = QHBoxLayout(req_widget)
req_layout.setContentsMargins(0, 0, 0, 0)
req_layout.setAlignment(Qt.AlignCenter)
req_check = QCheckBox()
req_check.setChecked(required)
req_check.stateChanged.connect(self._update_preview)
req_layout.addWidget(req_check)
self.table.setCellWidget(row, 3, req_widget)
# Delete button
btn_delete = QPushButton("Delete")
btn_delete.setMinimumHeight(36)
btn_delete.setStyleSheet("""
QPushButton {
background-color: #e53e3e;
color: white;
border: none;
border-radius: 4px;
padding: 4px 8px;
}
QPushButton:hover {
background-color: #c53030;
}
""")
btn_delete.clicked.connect(lambda: self._remove_field(row))
self.table.setCellWidget(row, 4, btn_delete)
self._update_preview()
def _remove_field(self, row: int):
"""Remove a field row."""
# Find the actual row (it may have shifted)
sender = self.sender()
for r in range(self.table.rowCount()):
if self.table.cellWidget(r, 4) == sender:
self.table.removeRow(r)
break
self._update_preview()
def _clear_all(self):
"""Clear all fields."""
self.table.setRowCount(0)
self._update_preview()
def _update_preview(self):
"""Update the JSON schema preview."""
schema = self.get_schema()
if schema:
self.preview.setPlainText(json.dumps(schema, indent=2))
else:
self.preview.setPlainText("(no fields defined)")
def _load_schema(self, schema: dict):
"""Load an existing schema into the builder."""
self.table.setRowCount(0)
if not schema or "properties" not in schema:
return
required = schema.get("required", [])
properties = schema.get("properties", {})
for name, prop in properties.items():
field_type = prop.get("type", "string")
description = prop.get("description", "")
is_required = name in required
self._add_field(name, field_type, description, is_required)
def get_schema(self) -> dict:
"""Build and return the JSON schema from the table."""
properties = {}
required = []
for row in range(self.table.rowCount()):
name_widget = self.table.cellWidget(row, 0)
type_widget = self.table.cellWidget(row, 1)
desc_widget = self.table.cellWidget(row, 2)
req_widget = self.table.cellWidget(row, 3)
if not name_widget:
continue
name = name_widget.text().strip()
if not name:
continue
field_type = type_widget.currentText() if type_widget else "string"
description = desc_widget.text().strip() if desc_widget else ""
# Get checkbox from the widget container
req_check = req_widget.findChild(QCheckBox) if req_widget else None
is_required = req_check.isChecked() if req_check else False
prop = {"type": field_type}
if description:
prop["description"] = description
# Add array items type
if field_type == "array":
prop["items"] = {"type": "string"}
properties[name] = prop
if is_required:
required.append(name)
if not properties:
return None
schema = {
"type": "object",
"properties": properties
}
if required:
schema["required"] = required
return schema
class PromptStepDialog(QDialog):
"""Dialog for editing prompt steps."""
def __init__(self, parent, step: PromptStep = None):
super().__init__(parent)
self.setWindowTitle("Edit Prompt Step" if step else "Add Prompt Step")
self.setMinimumSize(500, 400)
self.setMinimumSize(500, 500)
self._step = step
self._output_schema = step.output_schema if step else None
self._setup_ui()
if step:
@ -98,6 +346,26 @@ class PromptStepDialog(QDialog):
self.retries_label = QLabel("Max retries:")
form.addRow(self.retries_label, self.retries_spin)
# Output schema builder
schema_row = QHBoxLayout()
self.schema_status = QLabel("Default schema (output, reasoning)")
self.schema_status.setStyleSheet("color: #718096;")
schema_row.addWidget(self.schema_status)
schema_row.addStretch()
self.btn_edit_schema = QPushButton("Edit Schema...")
self.btn_edit_schema.clicked.connect(self._edit_schema)
schema_row.addWidget(self.btn_edit_schema)
self.btn_clear_schema = QPushButton("Reset")
self.btn_clear_schema.setToolTip("Reset to default schema")
self.btn_clear_schema.clicked.connect(self._clear_schema)
self.btn_clear_schema.setVisible(False)
schema_row.addWidget(self.btn_clear_schema)
schema_widget = QWidget()
schema_widget.setLayout(schema_row)
self.schema_label = QLabel("Output schema:")
form.addRow(self.schema_label, schema_widget)
layout.addLayout(form)
# Prompt text
@ -133,6 +401,38 @@ class PromptStepDialog(QDialog):
is_plain = bool(state)
self.retries_label.setVisible(not is_plain)
self.retries_spin.setVisible(not is_plain)
self.schema_label.setVisible(not is_plain)
self.schema_status.setVisible(not is_plain)
self.btn_edit_schema.setVisible(not is_plain)
self.btn_clear_schema.setVisible(not is_plain and self._output_schema is not None)
def _edit_schema(self):
"""Open the schema builder dialog."""
dialog = SchemaBuilderDialog(self, self._output_schema)
if dialog.exec():
self._output_schema = dialog.get_schema()
self._update_schema_status()
def _clear_schema(self):
"""Reset to default schema."""
self._output_schema = None
self._update_schema_status()
def _update_schema_status(self):
"""Update the schema status label."""
if self._output_schema:
fields = list(self._output_schema.get("properties", {}).keys())
if len(fields) <= 3:
fields_text = ", ".join(fields)
else:
fields_text = ", ".join(fields[:3]) + f" +{len(fields) - 3} more"
self.schema_status.setText(f"Custom schema ({fields_text})")
self.schema_status.setStyleSheet("color: #38a169;") # Green
self.btn_clear_schema.setVisible(True)
else:
self.schema_status.setText("Default schema (output, reasoning)")
self.schema_status.setStyleSheet("color: #718096;") # Gray
self.btn_clear_schema.setVisible(False)
def _load_step(self, step: PromptStep):
"""Load step data into form."""
@ -159,6 +459,8 @@ class PromptStepDialog(QDialog):
# Structured output fields
self.plain_text_check.setChecked(step.plain_text)
self.retries_spin.setValue(step.max_retries)
self._output_schema = step.output_schema
self._update_schema_status()
self._on_plain_text_changed(step.plain_text)
def _validate_and_accept(self):
@ -183,13 +485,17 @@ class PromptStepDialog(QDialog):
profile = None
# Get name, use None if empty
name = self.name_input.text().strip() or None
# Preserve prompt_file from original step if editing
prompt_file = self._step.prompt_file if self._step else None
return PromptStep(
prompt=self.prompt_input.toPlainText(),
provider=self.provider_combo.currentText(),
output_var=self.output_input.text().strip(),
prompt_file=prompt_file,
profile=profile,
name=name,
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()
)

View File

@ -58,11 +58,12 @@ class ProvidersPage(QWidget):
# Providers table
self.table = QTableWidget()
self.table.setColumnCount(3)
self.table.setHorizontalHeaderLabels(["Name", "Command", "Description"])
self.table.setColumnCount(4)
self.table.setHorizontalHeaderLabels(["Name", "Command", "Description", "Fallback"])
self.table.horizontalHeader().setSectionResizeMode(0, QHeaderView.ResizeToContents)
self.table.horizontalHeader().setSectionResizeMode(1, QHeaderView.Stretch)
self.table.horizontalHeader().setSectionResizeMode(2, QHeaderView.Stretch)
self.table.horizontalHeader().setSectionResizeMode(3, QHeaderView.ResizeToContents)
self.table.setSelectionBehavior(QTableWidget.SelectRows)
self.table.setSelectionMode(QTableWidget.SingleSelection)
self.table.verticalHeader().setVisible(False)
@ -107,6 +108,7 @@ class ProvidersPage(QWidget):
self.table.setItem(row, 0, name_item)
self.table.setItem(row, 1, QTableWidgetItem(provider.command))
self.table.setItem(row, 2, QTableWidgetItem(provider.description or ""))
self.table.setItem(row, 3, QTableWidgetItem(provider.fallback or ""))
self._on_selection_changed()

View File

@ -9,7 +9,7 @@ import yaml
from PySide6.QtWidgets import (
QWidget, QVBoxLayout, QHBoxLayout, QSplitter,
QTreeWidget, QTreeWidgetItem, QTextEdit, QLabel,
QPushButton, QGroupBox, QMessageBox, QFrame
QPushButton, QGroupBox, QMessageBox, QFrame, QLineEdit
)
from PySide6.QtCore import Qt, QThread, Signal, QTimer
from PySide6.QtGui import QFont, QColor, QBrush, QShortcut, QKeySequence
@ -402,6 +402,13 @@ class ToolsPage(QWidget):
left_layout.setContentsMargins(0, 0, 0, 0)
left_layout.setSpacing(8)
# Search box
self.search_box = QLineEdit()
self.search_box.setPlaceholderText("Search tools...")
self.search_box.setClearButtonEnabled(True)
self.search_box.textChanged.connect(self._filter_tools)
left_layout.addWidget(self.search_box)
self.tool_tree = QTreeWidget()
self.tool_tree.setHeaderHidden(True)
self.tool_tree.setIndentation(16)
@ -555,6 +562,7 @@ class ToolsPage(QWidget):
def refresh(self):
"""Refresh the tool list."""
self.search_box.clear()
self.tool_tree.clear()
self._current_tool = None
self.info_text.clear()
@ -741,6 +749,43 @@ class ToolsPage(QWidget):
self.refresh()
self.main_window.show_status(f"Status updated for '{tool_name}'")
def _filter_tools(self, text: str):
"""Filter tools by name or description based on search text."""
search = text.lower().strip()
for i in range(self.tool_tree.topLevelItemCount()):
category_item = self.tool_tree.topLevelItem(i)
visible_children = 0
for j in range(category_item.childCount()):
tool_item = category_item.child(j)
tool_name = tool_item.data(0, Qt.UserRole)
if not search:
# No filter - show all
tool_item.setHidden(False)
visible_children += 1
else:
# Check if name matches
matches = search in tool_name.lower()
# Also check description if we have the tool loaded
if not matches:
tool = load_tool(tool_name)
if tool and tool.description:
matches = search in tool.description.lower()
tool_item.setHidden(not matches)
if matches:
visible_children += 1
# Hide category if no children match
category_item.setHidden(visible_children == 0)
# Expand categories when filtering to show results
if search and visible_children > 0:
category_item.setExpanded(True)
def _on_selection_changed(self):
"""Handle tool selection change."""
items = self.tool_tree.selectedItems()

View File

@ -21,13 +21,17 @@ class Provider:
name: str
command: str
description: str = ""
fallback: Optional[str] = None # Name of fallback provider if this one fails
def to_dict(self) -> dict:
return {
d = {
"name": self.name,
"command": self.command,
"description": self.description,
}
if self.fallback:
d["fallback"] = self.fallback
return d
@classmethod
def from_dict(cls, data: dict) -> "Provider":
@ -35,6 +39,7 @@ class Provider:
name=data["name"],
command=data["command"],
description=data.get("description", ""),
fallback=data.get("fallback"),
)
@ -139,7 +144,7 @@ def delete_provider(name: str) -> bool:
return False
def call_provider(provider_name: str, prompt: str, timeout: int = 300) -> ProviderResult:
def call_provider(provider_name: str, prompt: str, timeout: int = 300, _tried: Optional[set] = None) -> ProviderResult:
"""
Call an AI provider with the given prompt.
@ -147,10 +152,16 @@ def call_provider(provider_name: str, prompt: str, timeout: int = 300) -> Provid
provider_name: Name of the provider to use
prompt: The prompt to send
timeout: Maximum execution time in seconds
_tried: Internal set of already-tried providers (prevents infinite loops)
Returns:
ProviderResult with the response text or error
"""
# Track which providers we've tried to prevent infinite fallback loops
if _tried is None:
_tried = set()
_tried.add(provider_name)
# Handle mock provider specially
if provider_name.lower() == "mock":
return mock_provider(prompt)
@ -175,13 +186,19 @@ def call_provider(provider_name: str, prompt: str, timeout: int = 300) -> Provid
# shlex failed (unbalanced quotes, etc.) - fall back to simple split
base_cmd = cmd.split()[0]
# Helper to try fallback provider if available
def try_fallback(error_msg: str) -> ProviderResult:
if provider.fallback and provider.fallback not in _tried:
import sys
print(f"[fallback] {provider_name} failed, trying {provider.fallback}...", file=sys.stderr)
return call_provider(provider.fallback, prompt, timeout, _tried)
return ProviderResult(text="", success=False, error=error_msg)
# Expand ~ for the which check
base_cmd_expanded = os.path.expanduser(base_cmd)
if not shutil.which(base_cmd_expanded) and not os.path.isfile(base_cmd_expanded):
return ProviderResult(
text="",
success=False,
error=f"Command '{base_cmd}' not found. Is it installed and in PATH?\n\nTo install AI providers, run: cmdforge providers install"
return try_fallback(
f"Command '{base_cmd}' not found. Is it installed and in PATH?\n\nTo install AI providers, run: cmdforge providers install"
)
try:
@ -198,11 +215,7 @@ def call_provider(provider_name: str, prompt: str, timeout: int = 300) -> Provid
error_msg = f"Provider exited with code {result.returncode}: {result.stderr}"
if "not found" in result.stderr.lower() or "not installed" in result.stderr.lower():
error_msg += "\n\nTo install AI providers, run: cmdforge providers install"
return ProviderResult(
text="",
success=False,
error=error_msg
)
return try_fallback(error_msg)
# Warn if output is empty (provider ran but returned nothing)
if not result.stdout.strip():
@ -217,57 +230,52 @@ def call_provider(provider_name: str, prompt: str, timeout: int = 300) -> Provid
provider_id = provider_match.group(1) if provider_match else "unknown"
model_id = model_match.group(1) if model_match else "unknown"
return ProviderResult(
text="",
success=False,
error=f"Model '{model_id}' from provider '{provider_id}' is not available.\n\n"
f"To fix this, either:\n"
f" 1. Run 'opencode' to connect the {provider_id} provider\n"
f" 2. Use --provider to pick a different model (e.g., --provider opencode-pickle)\n"
f" 3. Run 'cmdforge ui' to edit the tool's default provider"
return try_fallback(
f"Model '{model_id}' from provider '{provider_id}' is not available.\n\n"
f"To fix this, either:\n"
f" 1. Run 'opencode' to connect the {provider_id} provider\n"
f" 2. Use --provider to pick a different model (e.g., --provider opencode-pickle)\n"
f" 3. Run 'cmdforge ui' to edit the tool's default provider"
)
stderr_hint = f" (stderr: {stderr[:200]}...)" if len(stderr) > 200 else (f" (stderr: {stderr})" if stderr else "")
return ProviderResult(
text="",
success=False,
error=f"Provider returned empty output{stderr_hint}.\n\nThis may mean the model is not available. Try a different provider or run: cmdforge providers install"
return try_fallback(
f"Provider returned empty output{stderr_hint}.\n\nThis may mean the model is not available. Try a different provider or run: cmdforge providers install"
)
return ProviderResult(text=result.stdout, success=True)
except subprocess.TimeoutExpired:
return ProviderResult(
text="",
success=False,
error=f"Provider timed out after {timeout} seconds"
)
return try_fallback(f"Provider timed out after {timeout} seconds")
except Exception as e:
return ProviderResult(
text="",
success=False,
error=f"Provider error: {str(e)}"
)
return try_fallback(f"Provider error: {str(e)}")
def mock_provider(prompt: str) -> ProviderResult:
"""
Return a mock response for testing.
Returns structured JSON that matches the default output schema,
allowing tools to be tested without real AI providers.
Args:
prompt: The prompt (used for generating mock response)
Returns:
ProviderResult with mock response
ProviderResult with mock JSON response
"""
import json
lines = prompt.strip().split('\n')
preview = lines[0][:50] + "..." if len(lines[0]) > 50 else lines[0]
# Return structured JSON matching default schema
mock_response = {
"output": f"[MOCK] {preview}",
"reasoning": f"Mock response for prompt with {len(prompt)} chars, {len(lines)} lines."
}
return ProviderResult(
text=f"[MOCK RESPONSE]\n"
f"Prompt length: {len(prompt)} chars, {len(lines)} lines\n"
f"First line: {preview}\n"
f"\n"
f"This is a mock response. Use a real provider for actual output.",
text=json.dumps(mock_response),
success=True
)

View File

@ -2,6 +2,7 @@
import argparse
import json
import re
import sys
from pathlib import Path
from typing import Optional
@ -63,16 +64,110 @@ def validate_schema(data: dict, schema: dict) -> None:
print("Warning: jsonschema not installed, skipping schema validation", file=sys.stderr)
def append_schema_instructions(prompt: str, schema: dict) -> str:
def generate_schema_example(schema: dict) -> str:
"""
Append schema instructions to a prompt.
Generate a simple example JSON from a schema.
Small models work better with concrete examples than abstract JSON schemas.
Args:
schema: JSON schema to generate example from
Returns:
Example JSON string
"""
def build_example(schema_part: dict) -> any:
"""Recursively build example from schema."""
type_ = schema_part.get("type", "string")
if type_ == "string":
if "enum" in schema_part:
return schema_part["enum"][0]
desc = schema_part.get("description", "value")
return f"<{desc}>"
elif type_ == "number":
return 0.0
elif type_ == "integer":
return 0
elif type_ == "boolean":
return True
elif type_ == "array":
items = schema_part.get("items", {})
if items.get("type") == "object" and "properties" in items:
# Nested object array - generate one example item
return [build_example(items)]
elif items.get("type") == "string":
return ["<item1>", "<item2>"]
else:
return []
elif type_ == "object":
if "properties" in schema_part:
obj = {}
for key, spec in schema_part["properties"].items():
obj[key] = build_example(spec)
return obj
return {"...": "..."}
else:
return f"<{type_}>"
# Handle top-level schema
if schema.get("type") == "array":
example = build_example(schema)
elif schema.get("type") == "object" or "properties" in schema:
example = {}
for key, spec in schema.get("properties", {}).items():
example[key] = build_example(spec)
else:
example = build_example(schema)
return json.dumps(example, indent=2)
def prepend_schema_instructions(prompt: str, schema: dict) -> str:
"""
Prepend schema instructions to a prompt using example-based format.
Small models work better when format instructions come FIRST and use
concrete examples rather than abstract JSON schemas. This prevents
models from echoing the schema as content.
Args:
prompt: Original prompt text
schema: JSON schema the response must match
Returns:
Augmented prompt with schema instructions
Augmented prompt with schema instructions prepended
"""
example = generate_schema_example(schema)
# Build field guidance based on schema
props = schema.get("properties", {})
required = schema.get("required", [])
field_notes = []
if "reasoning" in props and "output" in props:
field_notes.append("Put your thinking in 'reasoning' and your answer in 'output'.")
if required:
field_notes.append(f"Required: {', '.join(required)}.")
field_guidance = " ".join(field_notes)
return f"""Respond with ONLY valid JSON in this format:
{example}
No text before or after the JSON. No markdown fences.{' ' + field_guidance if field_guidance else ''}
---
{prompt}"""
# Legacy function kept for rollback if needed
def _append_schema_instructions_legacy(prompt: str, schema: dict) -> str:
"""
DEPRECATED: Append schema instructions to a prompt.
This approach causes small models to echo the schema. Use
prepend_schema_instructions() instead.
"""
schema_json = json.dumps(schema, indent=2)
@ -323,6 +418,94 @@ def substitute_variables(template: str, variables: dict, warn_non_scalar: bool =
return result
def _extract_json(text: str) -> dict | list | None:
"""
Extract JSON from LLM response using multiple strategies.
LLMs often wrap JSON in markdown, add explanations before/after,
or include other text. This function tries several approaches
to find and extract valid JSON.
Returns parsed JSON (dict or list) or None if extraction fails.
"""
if not text:
return None
text = text.strip()
# Strategy 1: Direct parse (already clean JSON)
try:
return json.loads(text)
except json.JSONDecodeError:
pass
# Strategy 2: Strip markdown code fences
if text.startswith("```"):
cleaned = re.sub(r'^```\w*\n?', '', text)
cleaned = re.sub(r'\n?```\s*$', '', cleaned)
try:
return json.loads(cleaned.strip())
except json.JSONDecodeError:
pass
# Strategy 3: Find first { or [ and match to last } or ]
# This handles "Here's the JSON:\n{...}" or "{...}\nExplanation"
first_brace = text.find('{')
first_bracket = text.find('[')
if first_brace == -1 and first_bracket == -1:
return None
# Determine which comes first and what we're looking for
if first_bracket == -1 or (first_brace != -1 and first_brace < first_bracket):
start_char, end_char = '{', '}'
start_pos = first_brace
else:
start_char, end_char = '[', ']'
start_pos = first_bracket
# Find matching end by counting braces/brackets
depth = 0
end_pos = -1
in_string = False
escape_next = False
for i in range(start_pos, len(text)):
char = text[i]
if escape_next:
escape_next = False
continue
if char == '\\':
escape_next = True
continue
if char == '"' and not escape_next:
in_string = not in_string
continue
if in_string:
continue
if char == start_char:
depth += 1
elif char == end_char:
depth -= 1
if depth == 0:
end_pos = i
break
if end_pos != -1:
candidate = text[start_pos:end_pos + 1]
try:
return json.loads(candidate)
except json.JSONDecodeError:
pass
return None
def execute_prompt_step(
step: PromptStep,
variables: dict,
@ -379,8 +562,8 @@ def execute_prompt_step(
# Structured output mode - enforce JSON with schema validation
schema = step.output_schema or DEFAULT_OUTPUT_SCHEMA
# Augment prompt with schema instructions
augmented_prompt = append_schema_instructions(prompt, schema)
# Augment prompt with schema instructions (prepended for small model compatibility)
augmented_prompt = prepend_schema_instructions(prompt, schema)
max_attempts = step.max_retries + 1
last_response = None
@ -415,17 +598,10 @@ Please try again with valid JSON matching the schema exactly."""
last_response = result.text
# Strip markdown code fences if present (common model behavior)
text = result.text.strip()
if text.startswith("```"):
text = re.sub(r'^```\w*\n', '', text)
text = re.sub(r'\n```\s*$', '', text)
# Try to parse as JSON
try:
parsed = json.loads(text)
except json.JSONDecodeError as e:
last_error = f"Invalid JSON: {e}"
# Extract and parse JSON with multiple strategies
parsed = _extract_json(result.text)
if parsed is None:
last_error = f"Invalid JSON: Could not extract valid JSON from response"
if attempt == max_attempts - 1:
print(f"Prompt step failed after {step.max_retries} retry(s): {last_error}", file=sys.stderr)
if verbose:
@ -445,8 +621,9 @@ Please try again with valid JSON matching the schema exactly."""
return "", False
continue
# Success - return normalized JSON string
return json.dumps(parsed), True
# Success - return parsed dict for code step compatibility
# Template substitution handles dicts via _get_nested_value
return parsed, True
# Should not reach here, but just in case
return "", False

View File

@ -94,12 +94,18 @@ class TestMockProvider:
def test_mock_returns_success(self):
result = mock_provider("Test prompt")
assert result.success is True
assert "[MOCK RESPONSE]" in result.text
# Mock returns structured JSON with output and reasoning
import json
parsed = json.loads(result.text)
assert "[MOCK]" in parsed["output"]
def test_mock_includes_prompt_info(self):
result = mock_provider("This is a test prompt")
assert "Prompt length:" in result.text
assert "chars" in result.text
import json
parsed = json.loads(result.text)
# Prompt info is in reasoning field
assert "chars" in parsed["reasoning"]
assert "lines" in parsed["reasoning"]
def test_mock_shows_first_line_preview(self):
result = mock_provider("First line here\nSecond line")
@ -228,7 +234,10 @@ class TestCallProvider:
"""Mock provider should work without subprocess."""
result = call_provider("mock", "Test prompt")
assert result.success is True
assert "[MOCK RESPONSE]" in result.text
# Mock returns structured JSON
import json
parsed = json.loads(result.text)
assert "[MOCK]" in parsed["output"]
def test_call_nonexistent_provider(self, temp_providers_file):
save_providers([])
@ -376,3 +385,125 @@ class TestProviderCommandParsing:
# Error should show expanded path
home = os.environ.get("HOME", "")
assert "$HOME" not in result.error or home in result.error
class TestProviderFallback:
"""Tests for provider fallback functionality."""
@pytest.fixture
def temp_providers_file(self, tmp_path):
providers_file = tmp_path / ".cmdforge" / "providers.yaml"
with patch('cmdforge.providers.PROVIDERS_FILE', providers_file):
yield providers_file
def test_provider_with_fallback_to_dict(self):
"""Provider with fallback should include it in dict."""
provider = Provider(
name="cloud",
command="cloud-ai",
description="Cloud AI",
fallback="local"
)
d = provider.to_dict()
assert d["fallback"] == "local"
def test_provider_without_fallback_to_dict(self):
"""Provider without fallback should not include fallback key."""
provider = Provider(name="simple", command="simple-ai")
d = provider.to_dict()
assert "fallback" not in d
def test_provider_from_dict_with_fallback(self):
"""Provider.from_dict should handle fallback field."""
data = {
"name": "cloud",
"command": "cloud-ai",
"fallback": "local"
}
provider = Provider.from_dict(data)
assert provider.fallback == "local"
def test_provider_from_dict_without_fallback(self):
"""Provider.from_dict should handle missing fallback."""
data = {"name": "simple", "command": "simple-ai"}
provider = Provider.from_dict(data)
assert provider.fallback is None
def test_fallback_triggers_on_failure(self, temp_providers_file):
"""When primary provider fails, fallback should be tried."""
save_providers([
Provider("primary", "nonexistent-cmd", fallback="mock"),
Provider("mock", "mock")
])
result = call_provider("primary", "Test prompt")
# Fallback to mock should succeed
assert result.success is True
assert "[MOCK]" in result.text
def test_fallback_chain(self, temp_providers_file):
"""Fallback can chain to another provider with fallback."""
save_providers([
Provider("first", "nonexistent1", fallback="second"),
Provider("second", "nonexistent2", fallback="mock"),
Provider("mock", "mock")
])
result = call_provider("first", "Test prompt")
# Should chain: first -> second -> mock
assert result.success is True
assert "[MOCK]" in result.text
def test_fallback_prevents_infinite_loop(self, temp_providers_file):
"""Circular fallback references should not cause infinite loop."""
save_providers([
Provider("a", "nonexistent-a", fallback="b"),
Provider("b", "nonexistent-b", fallback="a"),
])
result = call_provider("a", "Test prompt")
# Should fail gracefully, not loop forever
assert result.success is False
assert "not found" in result.error.lower()
@patch('subprocess.run')
@patch('shutil.which')
def test_fallback_on_timeout(self, mock_which, mock_run, temp_providers_file):
"""Fallback should trigger on timeout."""
import subprocess
mock_which.return_value = "/usr/bin/slow-ai"
mock_run.side_effect = subprocess.TimeoutExpired("slow-ai", 5)
save_providers([
Provider("slow", "slow-ai", fallback="mock"),
Provider("mock", "mock")
])
result = call_provider("slow", "Test prompt", timeout=5)
assert result.success is True
assert "[MOCK]" in result.text
@patch('subprocess.run')
@patch('shutil.which')
def test_fallback_on_nonzero_exit(self, mock_which, mock_run, temp_providers_file):
"""Fallback should trigger on non-zero exit code."""
mock_which.return_value = "/usr/bin/failing-ai"
mock_run.return_value = MagicMock(
returncode=1,
stdout="",
stderr="API error"
)
save_providers([
Provider("failing", "failing-ai", fallback="mock"),
Provider("mock", "mock")
])
result = call_provider("failing", "Test prompt")
assert result.success is True
assert "[MOCK]" in result.text

View File

@ -268,11 +268,10 @@ class TestStructuredOutput:
output, success = execute_prompt_step(step, {"input": ""})
assert success is True
# Output should be normalized JSON
import json
parsed = json.loads(output)
assert parsed["output"] == "Hello, world!"
assert parsed["reasoning"] == "Greeting requested"
# Output should be a parsed dict (for code step compatibility)
assert isinstance(output, dict)
assert output["output"] == "Hello, world!"
assert output["reasoning"] == "Greeting requested"
@patch('cmdforge.runner.call_provider')
def test_structured_output_strips_markdown_fences(self, mock_call):
@ -291,9 +290,8 @@ class TestStructuredOutput:
output, success = execute_prompt_step(step, {"input": ""})
assert success is True
import json
parsed = json.loads(output)
assert parsed["output"] == "Hello"
assert isinstance(output, dict)
assert output["output"] == "Hello"
@patch('cmdforge.runner.call_provider')
def test_structured_output_retry_on_invalid_json(self, mock_call):
@ -366,9 +364,8 @@ class TestStructuredOutput:
output, success = execute_prompt_step(step, {"input": ""})
assert success is True
import json
parsed = json.loads(output)
assert parsed["score"] == 0.8
assert isinstance(output, dict)
assert output["score"] == 0.8
@patch('cmdforge.runner.call_provider')
def test_structured_output_schema_violation_triggers_retry(self, mock_call):
@ -414,6 +411,105 @@ class TestStructuredOutput:
assert mock_call.call_count == 1
class TestSchemaInstructions:
"""Tests for schema instruction generation."""
def test_generate_schema_example_default(self):
"""Default schema should generate output/reasoning example."""
from cmdforge.runner import generate_schema_example, DEFAULT_OUTPUT_SCHEMA
import json
example = generate_schema_example(DEFAULT_OUTPUT_SCHEMA)
parsed = json.loads(example)
assert "output" in parsed
assert "reasoning" in parsed
assert "<" in parsed["output"] # Placeholder format
def test_generate_schema_example_with_enum(self):
"""Enum fields should use first enum value."""
from cmdforge.runner import generate_schema_example
import json
schema = {
"type": "object",
"properties": {
"intent": {"type": "string", "enum": ["question", "task", "chat"]}
}
}
example = generate_schema_example(schema)
parsed = json.loads(example)
assert parsed["intent"] == "question" # First enum value
def test_generate_schema_example_with_array(self):
"""Array fields should generate array example."""
from cmdforge.runner import generate_schema_example
import json
schema = {
"type": "object",
"properties": {
"items": {"type": "array", "items": {"type": "string"}}
}
}
example = generate_schema_example(schema)
parsed = json.loads(example)
assert isinstance(parsed["items"], list)
assert len(parsed["items"]) == 2
def test_generate_schema_example_nested_object(self):
"""Nested objects should be recursively generated."""
from cmdforge.runner import generate_schema_example
import json
schema = {
"type": "object",
"properties": {
"scores": {
"type": "object",
"properties": {
"accuracy": {"type": "number"},
"completeness": {"type": "number"}
}
}
}
}
example = generate_schema_example(schema)
parsed = json.loads(example)
assert isinstance(parsed["scores"], dict)
assert parsed["scores"]["accuracy"] == 0.0
assert parsed["scores"]["completeness"] == 0.0
def test_prepend_schema_instructions_format(self):
"""Instructions should be prepended, not appended."""
from cmdforge.runner import prepend_schema_instructions, DEFAULT_OUTPUT_SCHEMA
prompt = "What is 2+2?"
result = prepend_schema_instructions(prompt, DEFAULT_OUTPUT_SCHEMA)
# Instructions should come before the prompt
assert result.startswith("Respond with ONLY valid JSON")
assert prompt in result
# Prompt should be after the separator
assert result.index("---") < result.index(prompt)
def test_prepend_schema_instructions_includes_guidance(self):
"""Should include field guidance for output/reasoning schemas."""
from cmdforge.runner import prepend_schema_instructions, DEFAULT_OUTPUT_SCHEMA
result = prepend_schema_instructions("Test", DEFAULT_OUTPUT_SCHEMA)
assert "reasoning" in result.lower()
assert "output" in result.lower()
assert "Required:" in result
class TestExecuteCodeStep:
"""Tests for code step execution."""