diff --git a/CLAUDE.md b/CLAUDE.md index c5f0204..468d609 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -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. diff --git a/src/cmdforge/gui/dialogs/provider_dialog.py b/src/cmdforge/gui/dialogs/provider_dialog.py index 2af6ac8..8265c97 100644 --- a/src/cmdforge/gui/dialogs/provider_dialog.py +++ b/src/cmdforge/gui/dialogs/provider_dialog.py @@ -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 diff --git a/src/cmdforge/gui/dialogs/step_dialog.py b/src/cmdforge/gui/dialogs/step_dialog.py index 5bf54c2..36bd35d 100644 --- a/src/cmdforge/gui/dialogs/step_dialog.py +++ b/src/cmdforge/gui/dialogs/step_dialog.py @@ -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() ) diff --git a/src/cmdforge/gui/pages/providers_page.py b/src/cmdforge/gui/pages/providers_page.py index 8512bb6..5a23d8a 100644 --- a/src/cmdforge/gui/pages/providers_page.py +++ b/src/cmdforge/gui/pages/providers_page.py @@ -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() diff --git a/src/cmdforge/gui/pages/tools_page.py b/src/cmdforge/gui/pages/tools_page.py index 6c6c1f0..d4824b0 100644 --- a/src/cmdforge/gui/pages/tools_page.py +++ b/src/cmdforge/gui/pages/tools_page.py @@ -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() diff --git a/src/cmdforge/providers.py b/src/cmdforge/providers.py index 114a882..e251ba1 100644 --- a/src/cmdforge/providers.py +++ b/src/cmdforge/providers.py @@ -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 ) diff --git a/src/cmdforge/runner.py b/src/cmdforge/runner.py index 5af4c81..8b70109 100644 --- a/src/cmdforge/runner.py +++ b/src/cmdforge/runner.py @@ -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 ["", ""] + 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 diff --git a/tests/test_providers.py b/tests/test_providers.py index 88451ed..3db3a5a 100644 --- a/tests/test_providers.py +++ b/tests/test_providers.py @@ -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 diff --git a/tests/test_runner.py b/tests/test_runner.py index a164702..233885f 100644 --- a/tests/test_runner.py +++ b/tests/test_runner.py @@ -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."""