Merge pull request #22612 from dannon/agent-custom-tool-reflection

Add producer/critic reflection loops to CustomToolAgent
This commit is contained in:
Marius van den Beek
2026-05-22 09:06:08 +02:00
committed by GitHub
4 changed files with 698 additions and 149 deletions
+329 -135
View File
@@ -10,7 +10,11 @@ from typing import (
)
import yaml
from pydantic import ValidationError
from pydantic import (
BaseModel,
Field,
ValidationError,
)
from pydantic_ai import Agent
from pydantic_ai.exceptions import (
ModelHTTPError,
@@ -29,7 +33,6 @@ from .base import (
AgentResponse,
AgentType,
BaseGalaxyAgent,
extract_result_content,
extract_structured_output,
GalaxyAgentDependencies,
)
@@ -53,33 +56,98 @@ def _find_validation_error(exc: BaseException) -> Optional[ValidationError]:
return None
class CritiqueReport(BaseModel):
"""Structured critique returned by the LLM critic.
Issues are split between *clarity* (text the user reads -- description,
labels, help text) and *idiomaticity* (tool shape -- defaults, exposed
options, container choice). The producer is re-rolled only when
``should_refine`` is true, which the critic should reserve for issues
significant enough to be worth another model call.
"""
clarity_issues: list[str] = Field(default_factory=list)
idiomaticity_issues: list[str] = Field(default_factory=list)
should_refine: bool = False
summary: str = ""
class CustomToolAgent(BaseGalaxyAgent):
"""Agent that creates custom Galaxy tools using UserToolSource schema.
Requires a model with structured output support. If the configured model
doesn't support structured output, returns an error guiding the operator
to configure an appropriate model.
Reflection: ``UserToolSource``'s pydantic validators catch structural
issues at construction time. Two opt-in loops handle the remainder:
- **Validator-driven retry** (default on): if the producer's output
fails validation, the producer is re-called once with the structured
error list and asked to fix specifically those issues. Cap of one
retry.
- **Quality critic + refine** (default off): an LLM critic reviews the
validated tool for clarity / idiomaticity issues that pydantic can't
see. If the critic flags significant issues, the producer is
re-rolled once with the critique. Cap of one refine; if refinement
breaks validation, the original tool is kept.
Both loops are gated on per-deployment config under
``inference_services.custom_tool``: ``validator_retry_enabled`` and
``quality_critic_enabled``. Default to validator-only behavior --
operators turn the critic on when they're willing to pay for it.
"""
agent_type = AgentType.CUSTOM_TOOL
DEFAULT_MAX_TOKENS = 16384
def __init__(self, deps: GalaxyAgentDependencies):
super().__init__(deps)
self._critic_agent: Optional[Agent[GalaxyAgentDependencies, CritiqueReport]] = None
def _requires_structured_output(self) -> bool:
return True
def _create_agent(self) -> Agent[GalaxyAgentDependencies, Any]:
"""Create agent with UserToolSource as the output type."""
"""Create agent with UserToolSource as the output type.
Sets output_retries=0 because the agent's explicit reflection loop
owns the validation retry (to provide a better prompt).
"""
return Agent(
self._get_model(),
deps_type=GalaxyAgentDependencies,
output_type=UserToolSource,
system_prompt=self.get_system_prompt(),
output_retries=0,
)
def get_system_prompt(self) -> str:
prompt_path = Path(__file__).parent / "prompts" / "custom_tool_structured.md"
return prompt_path.read_text()
def _get_critic_system_prompt(self) -> str:
prompt_path = Path(__file__).parent / "prompts" / "custom_tool_critic.md"
return prompt_path.read_text()
def _get_critic_agent(self) -> Agent[GalaxyAgentDependencies, CritiqueReport]:
"""Lazily build the critic agent. Same model as the producer by default;
operators can override via ``inference_services.custom_tool.critic_model``."""
if self._critic_agent is None:
self._critic_agent = Agent(
self._get_model(),
deps_type=GalaxyAgentDependencies,
output_type=CritiqueReport,
system_prompt=self._get_critic_system_prompt(),
)
return self._critic_agent
def _validator_retry_enabled(self) -> bool:
return bool(self._get_agent_config("validator_retry_enabled", True))
def _quality_critic_enabled(self) -> bool:
return bool(self._get_agent_config("quality_critic_enabled", False))
async def process(self, query: str, context: Optional[dict[str, Any]] = None) -> AgentResponse:
validation_error = self._validate_query(query)
if validation_error:
@@ -87,60 +155,215 @@ class CustomToolAgent(BaseGalaxyAgent):
capability_error = self._validate_model_capabilities()
if capability_error:
return self._build_response(
content=capability_error,
confidence=ConfidenceLevel.LOW,
method="capability_check",
query=query,
suggestions=[
ActionSuggestion(
action_type=ActionType.CONTACT_SUPPORT,
description="Contact your Galaxy administrator to configure AI tool generation",
parameters={},
confidence=ConfidenceLevel.HIGH,
priority=1,
)
],
error="model_capability",
agent_data={"requires": "structured_output"},
)
return self._capability_error_response(capability_error, query)
try:
result = await self._run_with_retry(query)
produced = await self._produce_tool(query)
if produced is None:
return self._invalid_structured_output_response(query)
if isinstance(produced, list):
# _produce_tool returns the formatted error list when the
# producer hit a pydantic ValidationError. Retry once if
# enabled, otherwise surface the issues to the user.
if not self._validator_retry_enabled():
return self._validation_failed_response(produced, query)
log.info("CustomTool: model failed validation (%d issue(s)); retrying once", len(produced))
retried = await self._produce_tool(query, retry_errors=produced)
if retried is None:
return self._invalid_structured_output_response(query)
if isinstance(retried, list):
return self._validation_failed_response(retried, query)
tool, tool_yaml, result = retried
else:
tool, tool_yaml, result = produced
# Quality critic: only refine on significant issues, never re-critique.
if self._quality_critic_enabled():
critique = await self._run_critic(tool_yaml, query)
if critique is not None and critique.should_refine:
log.info(
"CustomTool: critic flagged %d clarity / %d idiomaticity issues; refining once",
len(critique.clarity_issues),
len(critique.idiomaticity_issues),
)
refined = await self._produce_tool(query, critique=critique, prior_yaml=tool_yaml)
if isinstance(refined, tuple):
tool, tool_yaml, result = refined
elif isinstance(refined, list):
log.warning(
"CustomTool: refinement broke validation (%d issue(s)); keeping pre-refine tool",
len(refined),
)
return self._success_response(tool, tool_yaml, result, query)
except (OSError, ValueError) as e:
log.error(f"Tool creation error: {e}")
return self._build_response(
content=f"Failed to create tool: {str(e)}\n\nPlease try again with clear requirements.",
confidence=ConfidenceLevel.LOW,
method="error",
query=query,
error=str(e),
)
except ModelHTTPError as e:
return self._handle_model_http_error(e, query)
except UnexpectedModelBehavior as e:
return self._handle_unexpected_model_behavior(e, query)
async def _produce_tool(
self,
query: str,
retry_errors: Optional[list[str]] = None,
critique: Optional[CritiqueReport] = None,
prior_yaml: Optional[str] = None,
) -> Optional[tuple[UserToolSource, str, Any] | list[str]]:
"""Run the producer agent. Returns (tool, yaml, raw_result), error list, or None.
``retry_errors`` and ``critique`` are mutually exclusive: each prepends a
structured "fix specifically these issues" preamble to the original query.
"""
prompt = self._build_producer_prompt(query, retry_errors=retry_errors, critique=critique, prior_yaml=prior_yaml)
try:
result = await self._run_with_retry(prompt)
tool = extract_structured_output(result, UserToolSource, log)
if tool is None:
content = extract_result_content(result)
return self._build_response(
content=f"The model did not generate a valid tool definition. Response:\n\n{content}",
confidence=ConfidenceLevel.LOW,
method="text_fallback",
result=result,
query=query,
error="invalid_structured_output",
)
return None
lint_errors = lint_user_tool_source(tool)
if lint_errors:
log.debug("CustomToolAgent lint failure: %s", lint_errors)
bullet_text = "\n".join(f"- {issue}" for issue in lint_errors)
return self._build_response(
content=(
"The model produced a tool definition, but it has problems "
"that need to be fixed before it can be saved:\n\n"
f"{bullet_text}"
),
confidence=ConfidenceLevel.LOW,
method="lint_error",
query=query,
error="lint_failed",
agent_data={"lint_errors": lint_errors},
)
return lint_errors
tool_dict = tool.model_dump(by_alias=True, exclude_none=True)
tool_yaml = yaml.dump(tool_dict, default_flow_style=False, sort_keys=False)
return tool, tool_yaml, result
except UnexpectedModelBehavior as e:
pydantic_error = _find_validation_error(e)
if pydantic_error:
return format_validation_errors(pydantic_error)
raise e
response_content = f"""I've created a custom Galaxy tool:
@staticmethod
def _build_producer_prompt(
query: str,
retry_errors: Optional[list[str]] = None,
critique: Optional[CritiqueReport] = None,
prior_yaml: Optional[str] = None,
) -> str:
if not retry_errors and not critique:
return query
sections: list[str] = []
if retry_errors:
error_list = "\n".join(f"- {e}" for e in retry_errors)
sections.append(
"The previous attempt at this tool definition had specific problems "
"that need to be fixed:\n\n"
f"{error_list}\n\n"
"Re-generate the tool definition fixing exactly those issues. Keep "
"the parts that were already correct."
)
elif critique:
issue_lines: list[str] = []
if critique.clarity_issues:
issue_lines.append("Clarity issues:")
issue_lines.extend(f"- {issue}" for issue in critique.clarity_issues)
if critique.idiomaticity_issues:
if issue_lines:
issue_lines.append("")
issue_lines.append("Idiomaticity issues:")
issue_lines.extend(f"- {issue}" for issue in critique.idiomaticity_issues)
sections.append(
"The previous attempt is structurally valid but a reviewer flagged "
"the following quality issues:\n\n" + "\n".join(issue_lines) + "\n\n"
"Re-generate the tool definition addressing those issues. Don't "
"change parts the reviewer didn't flag."
)
if prior_yaml:
sections.append("Previous tool YAML:\n\n```yaml\n" + prior_yaml + "```")
sections.append("Original request:\n\n" + query)
return "\n\n".join(sections)
async def _run_critic(self, tool_yaml: str, query: str) -> Optional[CritiqueReport]:
"""Run the quality critic. Returns None if the critic call fails."""
critic = self._get_critic_agent()
critic_prompt = (
"Original request:\n\n"
f"{query}\n\n"
"Tool definition produced (already structurally validated):\n\n"
f"```yaml\n{tool_yaml}```\n\n"
"Critique this tool for clarity and idiomaticity. Set should_refine "
"only if the issues are significant enough to be worth another model call."
)
try:
result = await critic.run(critic_prompt, deps=self.deps)
output = getattr(result, "output", None)
if isinstance(output, CritiqueReport):
return output
log.warning("CustomTool: critic returned non-CritiqueReport output (%r); skipping refine", type(output))
return None
except (OSError, ValueError, ModelHTTPError, UnexpectedModelBehavior) as e:
log.warning("CustomTool: critic call failed (%s); skipping refine", e)
return None
def _capability_error_response(self, message: str, query: str) -> AgentResponse:
return self._build_response(
content=message,
confidence=ConfidenceLevel.LOW,
method="capability_check",
query=query,
suggestions=[
ActionSuggestion(
action_type=ActionType.CONTACT_SUPPORT,
description="Contact your Galaxy administrator to configure AI tool generation",
parameters={},
confidence=ConfidenceLevel.HIGH,
priority=1,
)
],
error="model_capability",
agent_data={"requires": "structured_output"},
)
def _invalid_structured_output_response(self, query: str) -> AgentResponse:
return self._build_response(
content="The model did not generate a valid tool definition.",
confidence=ConfidenceLevel.LOW,
method="text_fallback",
query=query,
error="invalid_structured_output",
)
def _validation_failed_response(
self,
validation_errors: list[str],
query: str,
) -> AgentResponse:
log.warning(
"CustomToolAgent produced a UserToolSource that failed validation: %s",
validation_errors,
)
bullet_list = "\n".join(f"- {issue}" for issue in validation_errors)
content = (
"The model produced a tool definition, but it has problems "
"that need to be fixed before it can be saved:\n\n"
f"{bullet_list}"
)
return self._build_response(
content=content,
confidence=ConfidenceLevel.LOW,
method="validation_error",
query=query,
error="validation_failed",
agent_data={
"validation_errors": validation_errors,
},
)
def _success_response(self, tool: UserToolSource, tool_yaml: str, result: Any, query: str) -> AgentResponse:
response_content = f"""I've created a custom Galaxy tool:
```yaml
{tool_yaml}
@@ -153,98 +376,43 @@ class CustomToolAgent(BaseGalaxyAgent):
The tool is ready to be saved and used in Galaxy."""
suggestions = [
ActionSuggestion(
action_type=ActionType.SAVE_TOOL,
description="Save this tool to Galaxy",
parameters={"tool_yaml": tool_yaml, "tool_id": tool.id},
confidence=ConfidenceLevel.HIGH,
priority=1,
),
]
return self._build_response(
content=response_content,
suggestions = [
ActionSuggestion(
action_type=ActionType.SAVE_TOOL,
description="Save this tool to Galaxy",
parameters={"tool_yaml": tool_yaml, "tool_id": tool.id},
confidence=ConfidenceLevel.HIGH,
method="structured",
result=result,
query=query,
suggestions=suggestions,
agent_data={
"tool_id": tool.id,
"tool_name": tool.name,
"tool_yaml": tool_yaml,
},
)
priority=1,
),
]
except (OSError, ValueError) as e:
log.error(f"Tool creation error: {e}")
return self._build_response(
content=f"Failed to create tool: {str(e)}\n\nPlease try again with clear requirements.",
confidence=ConfidenceLevel.LOW,
method="error",
query=query,
error=str(e),
)
except ModelHTTPError as e:
# Schema/grammar errors from model backends (vLLM, LiteLLM, etc.)
error_str = str(e).lower()
if "grammar" in error_str or "$defs" in error_str or "pointer" in error_str:
log.warning(f"Tool creation schema error (model may not support complex JSON schemas): {e}")
model = self._get_agent_config("model", "unknown")
return self._build_response(
content=(
f"The model '{model}' failed to generate a tool definition due to JSON schema limitations. "
"This typically happens with local inference backends (vLLM, LiteLLM proxies) that don't "
"support complex nested JSON schemas.\n\n"
"To resolve this, configure a model that fully supports structured output "
"(e.g., gpt-4o, claude-3-sonnet) via their native APIs."
),
confidence=ConfidenceLevel.LOW,
method="error",
query=query,
suggestions=[
ActionSuggestion(
action_type=ActionType.CONTACT_SUPPORT,
description="Contact support for help configuring a compatible model",
parameters={},
confidence=ConfidenceLevel.HIGH,
priority=1,
)
],
error="schema_limitation",
agent_data={"model": model},
)
raise
except UnexpectedModelBehavior as e:
pydantic_error = _find_validation_error(e)
if pydantic_error is not None:
# Expected path: the user is editing the prompt iteratively and
# the LLM produced a tool definition that fails one of the
# UserToolSource validators. Surface the friendly bullet list
# rather than a generic "model misbehaved" message.
bullets = format_validation_errors(pydantic_error)
log.debug("CustomToolAgent validation failure: %s", bullets)
bullet_text = "\n".join(f"- {issue}" for issue in bullets)
return self._build_response(
content=(
"The model produced a tool definition, but it has problems "
"that need to be fixed before it can be saved:\n\n"
f"{bullet_text}"
),
confidence=ConfidenceLevel.LOW,
method="validation_error",
query=query,
error="validation_failed",
agent_data={"validation_errors": bullets},
)
log.warning(f"Model failed to produce valid tool definition: {e}")
return self._build_response(
content=response_content,
confidence=ConfidenceLevel.HIGH,
method="structured",
result=result,
query=query,
suggestions=suggestions,
agent_data={
"tool_id": tool.id,
"tool_name": tool.name,
"tool_yaml": tool_yaml,
},
)
def _handle_model_http_error(self, e: ModelHTTPError, query: str) -> AgentResponse:
# Schema/grammar errors from model backends (vLLM, LiteLLM, etc.)
error_str = str(e).lower()
if "grammar" in error_str or "$defs" in error_str or "pointer" in error_str:
log.warning(f"Tool creation schema error (model may not support complex JSON schemas): {e}")
model = self._get_agent_config("model", "unknown")
return self._build_response(
content=(
f"The model '{model}' was unable to generate a valid tool definition after multiple attempts. "
"This may indicate the model doesn't fully support the required structured output format.\n\n"
"Try using a model with better structured output support (e.g., gpt-4o, claude-3-sonnet)."
f"The model '{model}' failed to generate a tool definition due to JSON schema limitations. "
"This typically happens with local inference backends (vLLM, LiteLLM proxies) that don't "
"support complex nested JSON schemas.\n\n"
"To resolve this, configure a model that fully supports structured output "
"(e.g., gpt-4o, claude-3-sonnet) via their native APIs."
),
confidence=ConfidenceLevel.LOW,
method="error",
@@ -258,6 +426,32 @@ The tool is ready to be saved and used in Galaxy."""
priority=1,
)
],
error="validation_failure",
error="schema_limitation",
agent_data={"model": model},
)
raise e
def _handle_unexpected_model_behavior(self, e: UnexpectedModelBehavior, query: str) -> AgentResponse:
log.warning(f"Model failed to produce valid tool definition: {e}")
model = self._get_agent_config("model", "unknown")
return self._build_response(
content=(
f"The model '{model}' was unable to generate a valid tool definition after multiple attempts. "
"This may indicate the model doesn't fully support the required structured output format.\n\n"
"Try using a model with better structured output support (e.g., gpt-4o, claude-3-sonnet)."
),
confidence=ConfidenceLevel.LOW,
method="error",
query=query,
suggestions=[
ActionSuggestion(
action_type=ActionType.CONTACT_SUPPORT,
description="Contact support for help configuring a compatible model",
parameters={},
confidence=ConfidenceLevel.HIGH,
priority=1,
)
],
error="validation_failure",
agent_data={"model": model},
)
@@ -0,0 +1,40 @@
# Galaxy Custom Tool Critic
You are a senior reviewer of Galaxy tool definitions. Another model has produced a tool definition that already passed structural validation -- IDs are well-formed, all referenced inputs are declared, container shape is recognized, citations are present. Your job is the **fuzzy quality** pass that validation can't do: clarity, idiomaticity, sensible defaults, helpful text.
You receive the original user request, the produced tool YAML, and you return a structured critique.
## What to flag
**Clarity issues** -- text that an end user will read:
- `description` doesn't say what the tool actually does, or is too generic ("Run the tool", "Process input")
- `name` is opaque or doesn't match the description
- Input `label` text is missing or duplicates the parameter name
- Input `help` text is missing for non-obvious parameters
- Output `label` text is missing or unclear
**Idiomaticity issues** -- shape of the tool:
- `shell_command` mixes shell quoting that won't escape correctly (e.g., bare `$(date)` instead of `\$(date)`)
- Optional parameters have no `default`, forcing the user to supply values that should be sensible
- Common analysis options aren't exposed (e.g., a BWA tool with no `-t` threads input)
- File outputs declared without `from_work_dir` or matching command output (the validator should have caught these, but flag any borderline cases)
- Container is a generic image like `ubuntu:latest` when a biocontainer for the wrapped tool exists
## What NOT to flag
- Anything the deterministic validator already catches (undeclared `inputs.X` references, container shape, citations, tool id format) -- assume it passed
- Style preferences that don't affect correctness or clarity ("I'd name this differently")
- Suggestions that would require new inputs or a fundamentally different tool design -- you are reviewing what's there, not redesigning
## Output
Return a `CritiqueReport` with:
- `clarity_issues`: list of concrete fixable issues, one per item. Empty list if none.
- `idiomaticity_issues`: list of concrete fixable issues. Empty list if none.
- `should_refine`: true only if at least one issue is significant enough that re-rolling the tool is worth a model call. Cosmetic-only critiques should set this to false.
- `summary`: one sentence describing the overall verdict.
Be parsimonious. The producer will be re-called with your critique if `should_refine` is true, which costs another LLM call. Don't trigger refinement for trivial issues.
@@ -24,7 +24,7 @@ You are a Galaxy tool generator. Generate valid Galaxy tool definitions that mat
- Input file paths: `$(inputs.param_name.path)` for single files
- Input values: `$(inputs.param_name)` for text, integer, float, boolean
- For array inputs: `$(inputs.param_name[].path)`
- Outputs are captured via `from_work_dir` in output definitions
- Outputs are captured via `from_work_dir` or `discover_datasets` in output definitions
## Input Parameter Types
@@ -38,16 +38,17 @@ Each input must have a `type` field. Valid types:
- **select**: Dropdown with options
Example input:
```yaml
inputs:
- name: input_file
type: data
format: fastq
label: Input FASTQ file
- name: num_threads
type: integer
default: 4
label: Number of threads
- name: input_file
type: data
format: fastq
label: Input FASTQ file
- name: num_threads
type: integer
default: 4
label: Number of threads
```
## Output Types
@@ -58,13 +59,14 @@ Each output must have a `type` field. Common types:
- **collection**: Collection of output files
Example output:
```yaml
outputs:
- name: output_file
type: data
format: sam
from_work_dir: aligned.sam
label: Aligned reads
- name: output_file
type: data
format: sam
from_work_dir: aligned.sam
label: Aligned reads
```
## Important Guidelines
+313
View File
@@ -35,7 +35,9 @@ import pytest
# Skip entire module if pydantic_ai is not installed
pydantic_ai = pytest.importorskip("pydantic_ai")
from pydantic import ValidationError
from pydantic_ai import Agent
from pydantic_ai.exceptions import UnexpectedModelBehavior
from pydantic_ai.messages import (
ModelMessage,
ModelRequest,
@@ -57,6 +59,7 @@ from galaxy.agents import (
ToolRecommendationAgent,
)
from galaxy.agents.base import truncate_message_history
from galaxy.agents.custom_tool import CritiqueReport
from galaxy.agents.registry import build_default_registry
from galaxy.agents.tools import SimplifiedToolRecommendationResult
@@ -1366,3 +1369,313 @@ class TestAgentUnitLiveLLM:
assert response.metadata["method"] == "simple_template"
assert "tool_id" in response.metadata
assert "tool_yaml" in response.metadata
def _validation_error_for(payload: dict) -> ValidationError:
"""Trigger a real ValidationError from UserToolSource for use in mocks.
Returning a real ValidationError (rather than a hand-constructed one)
keeps the tests honest about what pydantic-ai actually surfaces when
the producer's structured output fails validation.
"""
try:
UserToolSource(**payload)
except ValidationError as e:
return e
raise AssertionError(f"Expected ValidationError for payload: {payload}")
def _producer_validation_failure() -> UnexpectedModelBehavior:
"""Build the exception pydantic-ai raises when output validation fails.
With ``output_retries=0`` on the producer Agent, pydantic-ai wraps the
underlying ``ValidationError`` in ``UnexpectedModelBehavior`` and
bubbles it up via ``__cause__`` -- which is what ``_find_validation_error``
walks.
"""
ve = _validation_error_for(
{
"class": "GalaxyUserTool",
"id": "Bad-ID-Caps", # capital letters fail the model regex
"name": "Bad",
"version": "0.1.0",
"container": "ubuntu:latest",
"shell_command": "echo hi",
"inputs": [],
"outputs": [],
}
)
exc = UnexpectedModelBehavior("output validation failed")
exc.__cause__ = ve
return exc
def _valid_tool(name: str = "Echo Tool") -> UserToolSource:
return UserToolSource(
**{
"class": "GalaxyUserTool",
"id": "echo-tool",
"name": name,
"version": "0.1.0",
"description": "echo input to a file",
"container": "quay.io/biocontainers/python:3.13",
"shell_command": "echo '$(inputs.message)' > out.txt",
"inputs": [
{"name": "message", "type": "text", "value": "hi"},
],
"outputs": [
{"name": "out", "type": "data", "format": "txt", "from_work_dir": "out.txt"},
],
"citations": [{"type": "doi", "content": "10.1093/bioinformatics/btx123"}],
}
)
def _mock_run_result(tool: UserToolSource) -> mock.Mock:
result = mock.Mock()
result.output = tool
return result
class TestCustomToolAgentReflection:
"""Tests for CustomToolAgent's validator-retry and quality-critic loops.
With #22615's integrated validation, pydantic-ai wraps any UserToolSource
ValidationError in UnexpectedModelBehavior; the agent's reflection logic
surfaces those as a structured retry prompt or low-confidence response.
"""
def setup_method(self):
self.mock_config = mock.Mock()
self.mock_config.ai_api_key = "test-key"
self.mock_config.ai_model = "gpt-4o"
self.mock_config.ai_api_base_url = "http://localhost:4000/v1/"
self.mock_config.inference_services = None
self.mock_user = mock.Mock()
self.mock_user.id = 1
self.mock_user.username = "test_user"
self.mock_trans = mock.Mock()
self.mock_trans.app.config = self.mock_config
self.mock_trans.user = self.mock_user
self.deps = GalaxyAgentDependencies(
trans=self.mock_trans,
user=self.mock_user,
config=self.mock_config,
get_agent=agent_registry.get_agent,
job_manager=None,
)
@pytest.mark.asyncio
async def test_validator_retry_recovers_when_second_attempt_passes(self):
"""First call fails validation, retry returns a valid tool -> success."""
agent = CustomToolAgent(self.deps)
with mock.patch.object(
agent.agent,
"run",
side_effect=[
_producer_validation_failure(),
_mock_run_result(_valid_tool()),
],
) as mock_run:
response = await agent.process("Create a tool")
assert mock_run.call_count == 2
assert response.confidence == ConfidenceLevel.HIGH
assert response.metadata.get("tool_id") == "echo-tool"
# Retry prompt embeds the formatted error list as guidance.
retry_prompt = mock_run.call_args_list[1][0][0]
assert "previous attempt" in retry_prompt.lower()
@pytest.mark.asyncio
async def test_validator_retry_exhausted_returns_validation_failure(self):
"""Both producer calls fail validation -> low-confidence validation_failed."""
agent = CustomToolAgent(self.deps)
with mock.patch.object(
agent.agent,
"run",
side_effect=[
_producer_validation_failure(),
_producer_validation_failure(),
],
) as mock_run:
response = await agent.process("Create a tool")
assert mock_run.call_count == 2
assert response.confidence == ConfidenceLevel.LOW
assert response.metadata.get("error") == "validation_failed"
@pytest.mark.asyncio
async def test_validator_retry_disabled_short_circuits(self):
"""With validator_retry_enabled=False, producer is called exactly once."""
self.mock_config.inference_services = {
"custom_tool": {"validator_retry_enabled": False},
}
agent = CustomToolAgent(self.deps)
with mock.patch.object(
agent.agent,
"run",
side_effect=[_producer_validation_failure()],
) as mock_run:
response = await agent.process("Create a tool")
assert mock_run.call_count == 1
assert response.confidence == ConfidenceLevel.LOW
assert response.metadata.get("error") == "validation_failed"
@pytest.mark.asyncio
async def test_critic_disabled_by_default_skips_critic_call(self):
"""Default config: producer succeeds, critic is never invoked."""
agent = CustomToolAgent(self.deps)
with mock.patch.object(agent.agent, "run", return_value=_mock_run_result(_valid_tool())):
with mock.patch.object(agent, "_run_critic", new_callable=mock.AsyncMock) as mock_critic:
response = await agent.process("Create a tool")
mock_critic.assert_not_called()
assert response.confidence == ConfidenceLevel.HIGH
@pytest.mark.asyncio
async def test_critic_enabled_no_refine_when_should_refine_false(self):
"""Critic enabled but says nothing significant -> producer not re-rolled."""
self.mock_config.inference_services = {
"custom_tool": {"quality_critic_enabled": True},
}
agent = CustomToolAgent(self.deps)
no_issues = CritiqueReport(should_refine=False, summary="looks fine")
with mock.patch.object(agent.agent, "run", return_value=_mock_run_result(_valid_tool())) as mock_run:
with mock.patch.object(agent, "_run_critic", new_callable=mock.AsyncMock, return_value=no_issues):
response = await agent.process("Create a tool")
assert mock_run.call_count == 1
assert response.confidence == ConfidenceLevel.HIGH
@pytest.mark.asyncio
async def test_critic_enabled_refine_replaces_tool(self):
"""Critic flags refine -> producer is re-rolled and refined tool is used."""
self.mock_config.inference_services = {
"custom_tool": {"quality_critic_enabled": True},
}
agent = CustomToolAgent(self.deps)
original = _valid_tool(name="Echo Tool")
refined = _valid_tool(name="Echo Tool (refined)")
critique = CritiqueReport(
clarity_issues=["help text is terse"],
should_refine=True,
summary="needs clearer help",
)
with mock.patch.object(
agent.agent,
"run",
side_effect=[_mock_run_result(original), _mock_run_result(refined)],
) as mock_run:
with mock.patch.object(agent, "_run_critic", new_callable=mock.AsyncMock, return_value=critique):
response = await agent.process("Create a tool")
assert mock_run.call_count == 2
assert response.confidence == ConfidenceLevel.HIGH
assert "refined" in response.metadata.get("tool_yaml", "").lower()
@pytest.mark.asyncio
async def test_critic_refine_keeps_original_when_refinement_breaks_validation(self):
"""If refinement fails validation, the original (valid) tool is preserved."""
self.mock_config.inference_services = {
"custom_tool": {"quality_critic_enabled": True},
}
agent = CustomToolAgent(self.deps)
original = _valid_tool(name="Echo Tool")
critique = CritiqueReport(
idiomaticity_issues=["use a tighter container"],
should_refine=True,
summary="container is too broad",
)
with mock.patch.object(
agent.agent,
"run",
side_effect=[_mock_run_result(original), _producer_validation_failure()],
):
with mock.patch.object(agent, "_run_critic", new_callable=mock.AsyncMock, return_value=critique):
response = await agent.process("Create a tool")
assert response.confidence == ConfidenceLevel.HIGH
assert response.metadata.get("tool_id") == "echo-tool"
@pytestmark_live_llm
class TestAgentConsistencyLiveLLM:
"""Test agents with a consistent set of questions.
With the new router architecture using output functions, the router
handles queries directly or hands off to specialists. We test that
responses are appropriate for each query type.
"""
TEST_QUERIES = [
# Tool creation queries - should trigger custom_tool handoff
("Create a simple line counting tool", "tool_creation"),
("Build a Galaxy tool that runs samtools sort", "tool_creation"),
("I need a wrapper for BWA-MEM", "tool_creation"),
# Error analysis queries - should trigger error_analysis handoff
("Why did my job fail with exit code 127?", "error_analysis"),
("Help me debug this memory error", "error_analysis"),
("What does 'command not found' mean?", "error_analysis"),
# General queries - should get direct response from router
("Hello", "direct"),
("Thank you", "direct"),
("What can you do?", "direct"),
("How do I run BWA in Galaxy?", "direct"),
]
@pytest.fixture
def live_deps(self):
mock_config = mock.Mock()
mock_config.ai_api_key = os.environ.get("GALAXY_AI_API_KEY", "test-key")
mock_config.ai_model = os.environ.get("GALAXY_AI_MODEL", "llama-4-scout")
mock_config.ai_api_base_url = os.environ.get("GALAXY_AI_API_BASE_URL", "http://localhost:4000/v1/")
mock_user = mock.Mock()
mock_user.id = 1
mock_user.username = "test_user"
mock_trans = mock.Mock()
mock_trans.app.config = mock_config
mock_trans.user = mock_user
return GalaxyAgentDependencies(
trans=mock_trans,
user=mock_user,
config=mock_config,
get_agent=agent_registry.get_agent,
job_manager=None,
)
@pytest.mark.asyncio
async def test_response_consistency_live(self, live_deps):
router = QueryRouterAgent(live_deps)
for query, _query_type in self.TEST_QUERIES:
response = await router.process(query)
# All queries should return a response
assert response.content is not None, f"Query '{query}' should return content"
assert len(response.content) > 0, f"Query '{query}' should have non-empty content"
assert response.agent_type == "router"
@pytest.mark.asyncio
@pytest.mark.parametrize("query,query_type", TEST_QUERIES)
async def test_individual_query_response_live(self, live_deps, query, query_type):
router = QueryRouterAgent(live_deps)
response = await router.process(query)
# Verify we get a substantive response
assert response.content is not None
assert len(response.content) > 0
assert response.agent_type == "router"