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:
parent
8581775002
commit
71c6358b9d
42
CLAUDE.md
42
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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,10 +230,8 @@ 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"
|
||||
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"
|
||||
|
|
@ -228,46 +239,43 @@ def call_provider(provider_name: str, prompt: str, timeout: int = 300) -> Provid
|
|||
)
|
||||
|
||||
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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
Loading…
Reference in New Issue