"""
Robust tool-call parser for small language models.
Extracts structured tool calls from free-form model output using
regex-based parsing with multiple fallback strategies. Designed for
models that cannot reliably produce valid JSON (e.g. Qwen3-0.6B).
The expected model output format::
Thought: I need to search the documents for information about X.
retrieve_documents
information about X
Or for a final answer::
Thought: I have enough information to answer.
The answer is Y.
If parsing fails entirely, the raw output is treated as a direct answer.
"""
from __future__ import annotations
import re
from src.agents.schemas import ToolCall
# ---------------------------------------------------------------------------
# Regex patterns
# ---------------------------------------------------------------------------
_TOOL_RE = re.compile(r"\s*(\w+)\s*", re.IGNORECASE)
_PARAM_RE = re.compile(
r'(.*?)',
re.IGNORECASE | re.DOTALL,
)
_ANSWER_RE = re.compile(r"(.*?)", re.IGNORECASE | re.DOTALL)
_THOUGHT_RE = re.compile(
r"(?:Thought|Reasoning|Think):\s*(.+?)(?=||$)",
re.IGNORECASE | re.DOTALL,
)
# Fallback: model uses key: value\n... instead of
# the canonical / format. Common with small models.
_ALT_TOOL_RE = re.compile(
r"<(\w+)>\s*(.*?)\s*\1>",
re.IGNORECASE | re.DOTALL,
)
_ALT_PARAM_RE = re.compile(r"(\w+)\s*:\s*(.+)", re.MULTILINE)
# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------
def parse_tool_call(text: str) -> tuple[str, ToolCall | None, str | None]:
"""Parse model output into *(thought, tool_call_or_none, final_answer_or_none)*.
Exactly one of *tool_call* or *final_answer* will be non-``None``.
If neither is parseable the raw text is returned as the final answer.
"""
thought = _extract_thought(text)
# Try final answer first — if the model says , we're done.
final_answer = _extract_final_answer(text)
if final_answer is not None:
return thought, None, final_answer
# Try tool call.
tool_call = _extract_tool_call(text)
if tool_call is not None:
return thought, tool_call, None
# Fallback: treat entire output as a direct answer.
return thought, None, text.strip()
# ---------------------------------------------------------------------------
# Internal helpers
# ---------------------------------------------------------------------------
def _extract_thought(text: str) -> str:
"""Extract the ``Thought:`` line, or everything before the first tag."""
m = _THOUGHT_RE.search(text)
if m:
return m.group(1).strip()
# Fallback: everything before the first < tag.
idx = text.find("<")
if idx > 0:
return text[:idx].strip()
return ""
def _extract_tool_call(text: str) -> ToolCall | None:
"""Try to extract tool calls, with fallback for alternative formats.
Primary format: ``name`` + ``v``
Fallback format: ``key: value`` (common with small models)
"""
# Primary format: name + tags
m = _TOOL_RE.search(text)
if m is not None:
tool_name = m.group(1)
arguments: dict[str, str] = {}
for pm in _PARAM_RE.finditer(text):
arguments[pm.group(1)] = pm.group(2).strip()
return ToolCall(tool_name=tool_name, arguments=arguments, raw_text=text)
# Fallback: key: value\n...
# e.g. query: latest PyTorch\nnum_results: 3
for am in _ALT_TOOL_RE.finditer(text):
tag_name = am.group(1)
# Skip known non-tool tags
if tag_name.lower() in ("answer", "think", "param"):
continue
body = am.group(2)
arguments = {}
for pm in _ALT_PARAM_RE.finditer(body):
key = pm.group(1).strip().strip('"').strip("'")
val = pm.group(2).strip().strip('"').strip("'")
arguments[key] = val
return ToolCall(tool_name=tag_name, arguments=arguments, raw_text=text)
return None
def _extract_final_answer(text: str) -> str | None:
"""Try to extract ``...`` content."""
m = _ANSWER_RE.search(text)
if m is None:
return None
return m.group(1).strip()
def coerce_param(value: str, expected_type: str):
"""Best-effort type coercion for tool parameters."""
if expected_type == "int":
try:
return int(value)
except (ValueError, TypeError):
return value
if expected_type == "float":
try:
return float(value)
except (ValueError, TypeError):
return value
return value