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)
|
# Run all unit tests (excluding integration tests that need a server)
|
||||||
pytest tests/ -m "not integration"
|
pytest tests/ -m "not integration"
|
||||||
|
|
||||||
# Run all tests with verbose output
|
|
||||||
pytest tests/ -v
|
|
||||||
|
|
||||||
# Run a specific test file
|
# Run a specific test file
|
||||||
pytest tests/test_runner.py -v
|
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 -v
|
||||||
|
pytest tests/test_runner.py::TestSubstituteVariables::test_simple_substitution -v
|
||||||
|
|
||||||
# Run with coverage
|
# Run with coverage
|
||||||
pytest tests/ --cov=cmdforge --cov-report=html
|
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
|
python -m cmdforge.registry.app # Start server first
|
||||||
pytest tests/test_registry_integration.py -v -m integration
|
pytest tests/test_registry_integration.py -v -m integration
|
||||||
|
|
||||||
# Run the CLI
|
# CLI entry points
|
||||||
python -m cmdforge.cli
|
cmdforge # Main CLI / GUI launcher
|
||||||
|
cf # Interactive tool picker (fzf-style)
|
||||||
# Launch the GUI
|
python -m cmdforge.cli # Alternative CLI invocation
|
||||||
cmdforge
|
|
||||||
```
|
```
|
||||||
|
|
||||||
## Architecture
|
## Architecture
|
||||||
|
|
@ -76,7 +73,10 @@ Tools are YAML configs with:
|
||||||
|
|
||||||
### Step Types
|
### 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)
|
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`
|
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]'"
|
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
|
## Web UI & Registry
|
||||||
|
|
||||||
|
|
@ -145,14 +145,22 @@ CMDFORGE_REGISTRY_DB=/path/to/db PORT=5050 python -m cmdforge.web.app
|
||||||
|
|
||||||
## Testing Conventions
|
## Testing Conventions
|
||||||
|
|
||||||
Tests use `pytest` with common fixtures:
|
Tests use `pytest` without a shared conftest.py. Common patterns:
|
||||||
- `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
|
|
||||||
|
|
||||||
Mocking strategy:
|
**Mocking strategy:**
|
||||||
- **File system**: Use `tmp_path` fixture and patch module-level paths
|
- **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`
|
- **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.
|
Integration tests are marked with `@pytest.mark.integration` and require a running registry server.
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -2,10 +2,10 @@
|
||||||
|
|
||||||
from PySide6.QtWidgets import (
|
from PySide6.QtWidgets import (
|
||||||
QDialog, QVBoxLayout, QFormLayout, QLineEdit,
|
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):
|
class ProviderDialog(QDialog):
|
||||||
|
|
@ -43,12 +43,20 @@ class ProviderDialog(QDialog):
|
||||||
self.desc_input.setPlaceholderText("Claude AI via claude-cli")
|
self.desc_input.setPlaceholderText("Claude AI via claude-cli")
|
||||||
form.addRow("Description:", self.desc_input)
|
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)
|
layout.addLayout(form)
|
||||||
|
|
||||||
# Help text
|
# Help text
|
||||||
help_text = QLabel(
|
help_text = QLabel(
|
||||||
"The command should accept input on stdin and output to stdout.\n"
|
"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.setStyleSheet("color: #718096; font-size: 11px;")
|
||||||
help_text.setWordWrap(True)
|
help_text.setWordWrap(True)
|
||||||
|
|
@ -77,6 +85,11 @@ class ProviderDialog(QDialog):
|
||||||
self.name_input.setEnabled(False) # Can't rename
|
self.name_input.setEnabled(False) # Can't rename
|
||||||
self.cmd_input.setText(provider.command)
|
self.cmd_input.setText(provider.command)
|
||||||
self.desc_input.setText(provider.description or "")
|
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):
|
def _save(self):
|
||||||
"""Save the provider."""
|
"""Save the provider."""
|
||||||
|
|
@ -91,10 +104,17 @@ class ProviderDialog(QDialog):
|
||||||
self.cmd_input.setFocus()
|
self.cmd_input.setFocus()
|
||||||
return
|
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:
|
try:
|
||||||
add_provider(name, command, description)
|
add_provider(Provider(name, command, description, fallback))
|
||||||
self.accept()
|
self.accept()
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
from PySide6.QtWidgets import QMessageBox
|
from PySide6.QtWidgets import QMessageBox
|
||||||
|
|
|
||||||
|
|
@ -1,12 +1,14 @@
|
||||||
"""Step editor dialogs."""
|
"""Step editor dialogs."""
|
||||||
|
|
||||||
import ast
|
import ast
|
||||||
|
import json
|
||||||
|
|
||||||
from PySide6.QtWidgets import (
|
from PySide6.QtWidgets import (
|
||||||
QDialog, QVBoxLayout, QFormLayout, QLineEdit,
|
QDialog, QVBoxLayout, QFormLayout, QLineEdit,
|
||||||
QComboBox, QPushButton, QHBoxLayout, QLabel,
|
QComboBox, QPushButton, QHBoxLayout, QLabel,
|
||||||
QPlainTextEdit, QSplitter, QGroupBox, QTextEdit, QMessageBox,
|
QPlainTextEdit, QSplitter, QGroupBox, QTextEdit, QMessageBox,
|
||||||
QCheckBox, QSpinBox
|
QCheckBox, QSpinBox, QTableWidget, QTableWidgetItem, QHeaderView,
|
||||||
|
QWidget, QAbstractItemView
|
||||||
)
|
)
|
||||||
from PySide6.QtCore import Qt, QThread, Signal
|
from PySide6.QtCore import Qt, QThread, Signal
|
||||||
|
|
||||||
|
|
@ -15,14 +17,260 @@ from ...providers import load_providers, call_provider
|
||||||
from ...profiles import list_profiles
|
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):
|
class PromptStepDialog(QDialog):
|
||||||
"""Dialog for editing prompt steps."""
|
"""Dialog for editing prompt steps."""
|
||||||
|
|
||||||
def __init__(self, parent, step: PromptStep = None):
|
def __init__(self, parent, step: PromptStep = None):
|
||||||
super().__init__(parent)
|
super().__init__(parent)
|
||||||
self.setWindowTitle("Edit Prompt Step" if step else "Add Prompt Step")
|
self.setWindowTitle("Edit Prompt Step" if step else "Add Prompt Step")
|
||||||
self.setMinimumSize(500, 400)
|
self.setMinimumSize(500, 500)
|
||||||
self._step = step
|
self._step = step
|
||||||
|
self._output_schema = step.output_schema if step else None
|
||||||
self._setup_ui()
|
self._setup_ui()
|
||||||
|
|
||||||
if step:
|
if step:
|
||||||
|
|
@ -98,6 +346,26 @@ class PromptStepDialog(QDialog):
|
||||||
self.retries_label = QLabel("Max retries:")
|
self.retries_label = QLabel("Max retries:")
|
||||||
form.addRow(self.retries_label, self.retries_spin)
|
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)
|
layout.addLayout(form)
|
||||||
|
|
||||||
# Prompt text
|
# Prompt text
|
||||||
|
|
@ -133,6 +401,38 @@ class PromptStepDialog(QDialog):
|
||||||
is_plain = bool(state)
|
is_plain = bool(state)
|
||||||
self.retries_label.setVisible(not is_plain)
|
self.retries_label.setVisible(not is_plain)
|
||||||
self.retries_spin.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):
|
def _load_step(self, step: PromptStep):
|
||||||
"""Load step data into form."""
|
"""Load step data into form."""
|
||||||
|
|
@ -159,6 +459,8 @@ class PromptStepDialog(QDialog):
|
||||||
# Structured output fields
|
# Structured output fields
|
||||||
self.plain_text_check.setChecked(step.plain_text)
|
self.plain_text_check.setChecked(step.plain_text)
|
||||||
self.retries_spin.setValue(step.max_retries)
|
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)
|
self._on_plain_text_changed(step.plain_text)
|
||||||
|
|
||||||
def _validate_and_accept(self):
|
def _validate_and_accept(self):
|
||||||
|
|
@ -183,13 +485,17 @@ class PromptStepDialog(QDialog):
|
||||||
profile = None
|
profile = None
|
||||||
# Get name, use None if empty
|
# Get name, use None if empty
|
||||||
name = self.name_input.text().strip() or None
|
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(
|
return PromptStep(
|
||||||
prompt=self.prompt_input.toPlainText(),
|
prompt=self.prompt_input.toPlainText(),
|
||||||
provider=self.provider_combo.currentText(),
|
provider=self.provider_combo.currentText(),
|
||||||
output_var=self.output_input.text().strip(),
|
output_var=self.output_input.text().strip(),
|
||||||
|
prompt_file=prompt_file,
|
||||||
profile=profile,
|
profile=profile,
|
||||||
name=name,
|
name=name,
|
||||||
strip_fences=self.strip_fences_check.isChecked(),
|
strip_fences=self.strip_fences_check.isChecked(),
|
||||||
|
output_schema=self._output_schema,
|
||||||
plain_text=self.plain_text_check.isChecked(),
|
plain_text=self.plain_text_check.isChecked(),
|
||||||
max_retries=self.retries_spin.value()
|
max_retries=self.retries_spin.value()
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -58,11 +58,12 @@ class ProvidersPage(QWidget):
|
||||||
|
|
||||||
# Providers table
|
# Providers table
|
||||||
self.table = QTableWidget()
|
self.table = QTableWidget()
|
||||||
self.table.setColumnCount(3)
|
self.table.setColumnCount(4)
|
||||||
self.table.setHorizontalHeaderLabels(["Name", "Command", "Description"])
|
self.table.setHorizontalHeaderLabels(["Name", "Command", "Description", "Fallback"])
|
||||||
self.table.horizontalHeader().setSectionResizeMode(0, QHeaderView.ResizeToContents)
|
self.table.horizontalHeader().setSectionResizeMode(0, QHeaderView.ResizeToContents)
|
||||||
self.table.horizontalHeader().setSectionResizeMode(1, QHeaderView.Stretch)
|
self.table.horizontalHeader().setSectionResizeMode(1, QHeaderView.Stretch)
|
||||||
self.table.horizontalHeader().setSectionResizeMode(2, QHeaderView.Stretch)
|
self.table.horizontalHeader().setSectionResizeMode(2, QHeaderView.Stretch)
|
||||||
|
self.table.horizontalHeader().setSectionResizeMode(3, QHeaderView.ResizeToContents)
|
||||||
self.table.setSelectionBehavior(QTableWidget.SelectRows)
|
self.table.setSelectionBehavior(QTableWidget.SelectRows)
|
||||||
self.table.setSelectionMode(QTableWidget.SingleSelection)
|
self.table.setSelectionMode(QTableWidget.SingleSelection)
|
||||||
self.table.verticalHeader().setVisible(False)
|
self.table.verticalHeader().setVisible(False)
|
||||||
|
|
@ -107,6 +108,7 @@ class ProvidersPage(QWidget):
|
||||||
self.table.setItem(row, 0, name_item)
|
self.table.setItem(row, 0, name_item)
|
||||||
self.table.setItem(row, 1, QTableWidgetItem(provider.command))
|
self.table.setItem(row, 1, QTableWidgetItem(provider.command))
|
||||||
self.table.setItem(row, 2, QTableWidgetItem(provider.description or ""))
|
self.table.setItem(row, 2, QTableWidgetItem(provider.description or ""))
|
||||||
|
self.table.setItem(row, 3, QTableWidgetItem(provider.fallback or ""))
|
||||||
|
|
||||||
self._on_selection_changed()
|
self._on_selection_changed()
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -9,7 +9,7 @@ import yaml
|
||||||
from PySide6.QtWidgets import (
|
from PySide6.QtWidgets import (
|
||||||
QWidget, QVBoxLayout, QHBoxLayout, QSplitter,
|
QWidget, QVBoxLayout, QHBoxLayout, QSplitter,
|
||||||
QTreeWidget, QTreeWidgetItem, QTextEdit, QLabel,
|
QTreeWidget, QTreeWidgetItem, QTextEdit, QLabel,
|
||||||
QPushButton, QGroupBox, QMessageBox, QFrame
|
QPushButton, QGroupBox, QMessageBox, QFrame, QLineEdit
|
||||||
)
|
)
|
||||||
from PySide6.QtCore import Qt, QThread, Signal, QTimer
|
from PySide6.QtCore import Qt, QThread, Signal, QTimer
|
||||||
from PySide6.QtGui import QFont, QColor, QBrush, QShortcut, QKeySequence
|
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.setContentsMargins(0, 0, 0, 0)
|
||||||
left_layout.setSpacing(8)
|
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 = QTreeWidget()
|
||||||
self.tool_tree.setHeaderHidden(True)
|
self.tool_tree.setHeaderHidden(True)
|
||||||
self.tool_tree.setIndentation(16)
|
self.tool_tree.setIndentation(16)
|
||||||
|
|
@ -555,6 +562,7 @@ class ToolsPage(QWidget):
|
||||||
|
|
||||||
def refresh(self):
|
def refresh(self):
|
||||||
"""Refresh the tool list."""
|
"""Refresh the tool list."""
|
||||||
|
self.search_box.clear()
|
||||||
self.tool_tree.clear()
|
self.tool_tree.clear()
|
||||||
self._current_tool = None
|
self._current_tool = None
|
||||||
self.info_text.clear()
|
self.info_text.clear()
|
||||||
|
|
@ -741,6 +749,43 @@ class ToolsPage(QWidget):
|
||||||
self.refresh()
|
self.refresh()
|
||||||
self.main_window.show_status(f"Status updated for '{tool_name}'")
|
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):
|
def _on_selection_changed(self):
|
||||||
"""Handle tool selection change."""
|
"""Handle tool selection change."""
|
||||||
items = self.tool_tree.selectedItems()
|
items = self.tool_tree.selectedItems()
|
||||||
|
|
|
||||||
|
|
@ -21,13 +21,17 @@ class Provider:
|
||||||
name: str
|
name: str
|
||||||
command: str
|
command: str
|
||||||
description: str = ""
|
description: str = ""
|
||||||
|
fallback: Optional[str] = None # Name of fallback provider if this one fails
|
||||||
|
|
||||||
def to_dict(self) -> dict:
|
def to_dict(self) -> dict:
|
||||||
return {
|
d = {
|
||||||
"name": self.name,
|
"name": self.name,
|
||||||
"command": self.command,
|
"command": self.command,
|
||||||
"description": self.description,
|
"description": self.description,
|
||||||
}
|
}
|
||||||
|
if self.fallback:
|
||||||
|
d["fallback"] = self.fallback
|
||||||
|
return d
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_dict(cls, data: dict) -> "Provider":
|
def from_dict(cls, data: dict) -> "Provider":
|
||||||
|
|
@ -35,6 +39,7 @@ class Provider:
|
||||||
name=data["name"],
|
name=data["name"],
|
||||||
command=data["command"],
|
command=data["command"],
|
||||||
description=data.get("description", ""),
|
description=data.get("description", ""),
|
||||||
|
fallback=data.get("fallback"),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -139,7 +144,7 @@ def delete_provider(name: str) -> bool:
|
||||||
return False
|
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.
|
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
|
provider_name: Name of the provider to use
|
||||||
prompt: The prompt to send
|
prompt: The prompt to send
|
||||||
timeout: Maximum execution time in seconds
|
timeout: Maximum execution time in seconds
|
||||||
|
_tried: Internal set of already-tried providers (prevents infinite loops)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
ProviderResult with the response text or error
|
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
|
# Handle mock provider specially
|
||||||
if provider_name.lower() == "mock":
|
if provider_name.lower() == "mock":
|
||||||
return mock_provider(prompt)
|
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
|
# shlex failed (unbalanced quotes, etc.) - fall back to simple split
|
||||||
base_cmd = cmd.split()[0]
|
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
|
# Expand ~ for the which check
|
||||||
base_cmd_expanded = os.path.expanduser(base_cmd)
|
base_cmd_expanded = os.path.expanduser(base_cmd)
|
||||||
if not shutil.which(base_cmd_expanded) and not os.path.isfile(base_cmd_expanded):
|
if not shutil.which(base_cmd_expanded) and not os.path.isfile(base_cmd_expanded):
|
||||||
return ProviderResult(
|
return try_fallback(
|
||||||
text="",
|
f"Command '{base_cmd}' not found. Is it installed and in PATH?\n\nTo install AI providers, run: cmdforge providers install"
|
||||||
success=False,
|
|
||||||
error=f"Command '{base_cmd}' not found. Is it installed and in PATH?\n\nTo install AI providers, run: cmdforge providers install"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
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}"
|
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():
|
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"
|
error_msg += "\n\nTo install AI providers, run: cmdforge providers install"
|
||||||
return ProviderResult(
|
return try_fallback(error_msg)
|
||||||
text="",
|
|
||||||
success=False,
|
|
||||||
error=error_msg
|
|
||||||
)
|
|
||||||
|
|
||||||
# Warn if output is empty (provider ran but returned nothing)
|
# Warn if output is empty (provider ran but returned nothing)
|
||||||
if not result.stdout.strip():
|
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"
|
provider_id = provider_match.group(1) if provider_match else "unknown"
|
||||||
model_id = model_match.group(1) if model_match else "unknown"
|
model_id = model_match.group(1) if model_match else "unknown"
|
||||||
|
|
||||||
return ProviderResult(
|
return try_fallback(
|
||||||
text="",
|
f"Model '{model_id}' from provider '{provider_id}' is not available.\n\n"
|
||||||
success=False,
|
f"To fix this, either:\n"
|
||||||
error=f"Model '{model_id}' from provider '{provider_id}' is not available.\n\n"
|
f" 1. Run 'opencode' to connect the {provider_id} provider\n"
|
||||||
f"To fix this, either:\n"
|
f" 2. Use --provider to pick a different model (e.g., --provider opencode-pickle)\n"
|
||||||
f" 1. Run 'opencode' to connect the {provider_id} provider\n"
|
f" 3. Run 'cmdforge ui' to edit the tool's default provider"
|
||||||
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 "")
|
stderr_hint = f" (stderr: {stderr[:200]}...)" if len(stderr) > 200 else (f" (stderr: {stderr})" if stderr else "")
|
||||||
return ProviderResult(
|
return try_fallback(
|
||||||
text="",
|
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"
|
||||||
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 ProviderResult(text=result.stdout, success=True)
|
return ProviderResult(text=result.stdout, success=True)
|
||||||
|
|
||||||
except subprocess.TimeoutExpired:
|
except subprocess.TimeoutExpired:
|
||||||
return ProviderResult(
|
return try_fallback(f"Provider timed out after {timeout} seconds")
|
||||||
text="",
|
|
||||||
success=False,
|
|
||||||
error=f"Provider timed out after {timeout} seconds"
|
|
||||||
)
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
return ProviderResult(
|
return try_fallback(f"Provider error: {str(e)}")
|
||||||
text="",
|
|
||||||
success=False,
|
|
||||||
error=f"Provider error: {str(e)}"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def mock_provider(prompt: str) -> ProviderResult:
|
def mock_provider(prompt: str) -> ProviderResult:
|
||||||
"""
|
"""
|
||||||
Return a mock response for testing.
|
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:
|
Args:
|
||||||
prompt: The prompt (used for generating mock response)
|
prompt: The prompt (used for generating mock response)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
ProviderResult with mock response
|
ProviderResult with mock JSON response
|
||||||
"""
|
"""
|
||||||
|
import json
|
||||||
|
|
||||||
lines = prompt.strip().split('\n')
|
lines = prompt.strip().split('\n')
|
||||||
preview = lines[0][:50] + "..." if len(lines[0]) > 50 else lines[0]
|
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(
|
return ProviderResult(
|
||||||
text=f"[MOCK RESPONSE]\n"
|
text=json.dumps(mock_response),
|
||||||
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.",
|
|
||||||
success=True
|
success=True
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -2,6 +2,7 @@
|
||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
import json
|
import json
|
||||||
|
import re
|
||||||
import sys
|
import sys
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Optional
|
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)
|
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:
|
Args:
|
||||||
prompt: Original prompt text
|
prompt: Original prompt text
|
||||||
schema: JSON schema the response must match
|
schema: JSON schema the response must match
|
||||||
|
|
||||||
Returns:
|
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)
|
schema_json = json.dumps(schema, indent=2)
|
||||||
|
|
||||||
|
|
@ -323,6 +418,94 @@ def substitute_variables(template: str, variables: dict, warn_non_scalar: bool =
|
||||||
return result
|
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(
|
def execute_prompt_step(
|
||||||
step: PromptStep,
|
step: PromptStep,
|
||||||
variables: dict,
|
variables: dict,
|
||||||
|
|
@ -379,8 +562,8 @@ def execute_prompt_step(
|
||||||
# Structured output mode - enforce JSON with schema validation
|
# Structured output mode - enforce JSON with schema validation
|
||||||
schema = step.output_schema or DEFAULT_OUTPUT_SCHEMA
|
schema = step.output_schema or DEFAULT_OUTPUT_SCHEMA
|
||||||
|
|
||||||
# Augment prompt with schema instructions
|
# Augment prompt with schema instructions (prepended for small model compatibility)
|
||||||
augmented_prompt = append_schema_instructions(prompt, schema)
|
augmented_prompt = prepend_schema_instructions(prompt, schema)
|
||||||
|
|
||||||
max_attempts = step.max_retries + 1
|
max_attempts = step.max_retries + 1
|
||||||
last_response = None
|
last_response = None
|
||||||
|
|
@ -415,17 +598,10 @@ Please try again with valid JSON matching the schema exactly."""
|
||||||
|
|
||||||
last_response = result.text
|
last_response = result.text
|
||||||
|
|
||||||
# Strip markdown code fences if present (common model behavior)
|
# Extract and parse JSON with multiple strategies
|
||||||
text = result.text.strip()
|
parsed = _extract_json(result.text)
|
||||||
if text.startswith("```"):
|
if parsed is None:
|
||||||
text = re.sub(r'^```\w*\n', '', text)
|
last_error = f"Invalid JSON: Could not extract valid JSON from response"
|
||||||
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}"
|
|
||||||
if attempt == max_attempts - 1:
|
if attempt == max_attempts - 1:
|
||||||
print(f"Prompt step failed after {step.max_retries} retry(s): {last_error}", file=sys.stderr)
|
print(f"Prompt step failed after {step.max_retries} retry(s): {last_error}", file=sys.stderr)
|
||||||
if verbose:
|
if verbose:
|
||||||
|
|
@ -445,8 +621,9 @@ Please try again with valid JSON matching the schema exactly."""
|
||||||
return "", False
|
return "", False
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# Success - return normalized JSON string
|
# Success - return parsed dict for code step compatibility
|
||||||
return json.dumps(parsed), True
|
# Template substitution handles dicts via _get_nested_value
|
||||||
|
return parsed, True
|
||||||
|
|
||||||
# Should not reach here, but just in case
|
# Should not reach here, but just in case
|
||||||
return "", False
|
return "", False
|
||||||
|
|
|
||||||
|
|
@ -94,12 +94,18 @@ class TestMockProvider:
|
||||||
def test_mock_returns_success(self):
|
def test_mock_returns_success(self):
|
||||||
result = mock_provider("Test prompt")
|
result = mock_provider("Test prompt")
|
||||||
assert result.success is True
|
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):
|
def test_mock_includes_prompt_info(self):
|
||||||
result = mock_provider("This is a test prompt")
|
result = mock_provider("This is a test prompt")
|
||||||
assert "Prompt length:" in result.text
|
import json
|
||||||
assert "chars" in result.text
|
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):
|
def test_mock_shows_first_line_preview(self):
|
||||||
result = mock_provider("First line here\nSecond line")
|
result = mock_provider("First line here\nSecond line")
|
||||||
|
|
@ -228,7 +234,10 @@ class TestCallProvider:
|
||||||
"""Mock provider should work without subprocess."""
|
"""Mock provider should work without subprocess."""
|
||||||
result = call_provider("mock", "Test prompt")
|
result = call_provider("mock", "Test prompt")
|
||||||
assert result.success is True
|
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):
|
def test_call_nonexistent_provider(self, temp_providers_file):
|
||||||
save_providers([])
|
save_providers([])
|
||||||
|
|
@ -376,3 +385,125 @@ class TestProviderCommandParsing:
|
||||||
# Error should show expanded path
|
# Error should show expanded path
|
||||||
home = os.environ.get("HOME", "")
|
home = os.environ.get("HOME", "")
|
||||||
assert "$HOME" not in result.error or home in result.error
|
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": ""})
|
output, success = execute_prompt_step(step, {"input": ""})
|
||||||
|
|
||||||
assert success is True
|
assert success is True
|
||||||
# Output should be normalized JSON
|
# Output should be a parsed dict (for code step compatibility)
|
||||||
import json
|
assert isinstance(output, dict)
|
||||||
parsed = json.loads(output)
|
assert output["output"] == "Hello, world!"
|
||||||
assert parsed["output"] == "Hello, world!"
|
assert output["reasoning"] == "Greeting requested"
|
||||||
assert parsed["reasoning"] == "Greeting requested"
|
|
||||||
|
|
||||||
@patch('cmdforge.runner.call_provider')
|
@patch('cmdforge.runner.call_provider')
|
||||||
def test_structured_output_strips_markdown_fences(self, mock_call):
|
def test_structured_output_strips_markdown_fences(self, mock_call):
|
||||||
|
|
@ -291,9 +290,8 @@ class TestStructuredOutput:
|
||||||
output, success = execute_prompt_step(step, {"input": ""})
|
output, success = execute_prompt_step(step, {"input": ""})
|
||||||
|
|
||||||
assert success is True
|
assert success is True
|
||||||
import json
|
assert isinstance(output, dict)
|
||||||
parsed = json.loads(output)
|
assert output["output"] == "Hello"
|
||||||
assert parsed["output"] == "Hello"
|
|
||||||
|
|
||||||
@patch('cmdforge.runner.call_provider')
|
@patch('cmdforge.runner.call_provider')
|
||||||
def test_structured_output_retry_on_invalid_json(self, mock_call):
|
def test_structured_output_retry_on_invalid_json(self, mock_call):
|
||||||
|
|
@ -366,9 +364,8 @@ class TestStructuredOutput:
|
||||||
output, success = execute_prompt_step(step, {"input": ""})
|
output, success = execute_prompt_step(step, {"input": ""})
|
||||||
|
|
||||||
assert success is True
|
assert success is True
|
||||||
import json
|
assert isinstance(output, dict)
|
||||||
parsed = json.loads(output)
|
assert output["score"] == 0.8
|
||||||
assert parsed["score"] == 0.8
|
|
||||||
|
|
||||||
@patch('cmdforge.runner.call_provider')
|
@patch('cmdforge.runner.call_provider')
|
||||||
def test_structured_output_schema_violation_triggers_retry(self, mock_call):
|
def test_structured_output_schema_violation_triggers_retry(self, mock_call):
|
||||||
|
|
@ -414,6 +411,105 @@ class TestStructuredOutput:
|
||||||
assert mock_call.call_count == 1
|
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:
|
class TestExecuteCodeStep:
|
||||||
"""Tests for code step execution."""
|
"""Tests for code step execution."""
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue