Skip to content
135 changes: 125 additions & 10 deletions mellea/backends/tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -471,23 +471,128 @@ def find_func(d: object) -> tuple[str | None, Mapping | None]:
return None, None


# Granite 4.2's XML tool-call format:
# <tool_call>\n<function=NAME>\n<parameter=KEY>\nVALUE\n</parameter>\n</function>\n</tool_call>
# The body may hold only parameter blocks, so a value can contain any text except
# `</parameter>` (unescapable) and a broken call can't run into the next one.
# `</tool_call>` is optional, for truncated output.
_XML_TOOL_CALL_RE = re.compile(
r"<tool_call>\s*<function=([^>\n]+)>\s*"
r"((?:<parameter=[^>\n]+>(?:(?!</parameter>).)*</parameter>\s*)*)"
r"</function>\s*(?:</tool_call>)?",
re.DOTALL,
)
_XML_PARAMETER_RE = re.compile(r"<parameter=([^>\n]+)>(.*?)</parameter>", re.DOTALL)


def _json_tool_call_spans(text: str) -> list[tuple[int, int]]:
"""Return the `(start, end)` offsets in `text` of the calls `json_extraction` would find."""
spans: list[tuple[int, int]] = []
# Raw text, unlike `json_extraction`'s, so strings may hold literal newlines.
decoder = json.JSONDecoder(strict=False)
index = text.find("{")
while index != -1:
try:
obj, end = decoder.raw_decode(text, index)
except json.JSONDecodeError:
index = text.find("{", index + 1)
continue
name, args = find_func(obj)
if name is not None and args is not None:
spans.append((index, end))
index = text.find("{", end)
return spans


def _parse_xml_tool_calls(llm_response: str) -> tuple[list[tuple[str, Mapping]], str]:
"""Extract XML-format tool calls, skipping any quoted in a JSON call's arguments.

Values are raw strings, minus the newline the template wraps each one in.

Returns:
The calls, and `llm_response` with them blanked out.
"""
matches = list(_XML_TOOL_CALL_RE.finditer(llm_response))
if not matches:
return [], llm_response
json_spans = _json_tool_call_spans(llm_response)
matches = [
m
for m in matches
if not any(start <= m.start() < end for start, end in json_spans)
]

calls: list[tuple[str, Mapping]] = []
for match in matches:
args = {
name.strip(): value.removeprefix("\n").removesuffix("\n")
for name, value in _XML_PARAMETER_RE.findall(match.group(2))
}
calls.append((match.group(1).strip(), args))

remainder = llm_response
for match in reversed(matches): # back to front, so earlier offsets stay valid
remainder = f"{remainder[: match.start()]} {remainder[match.end() :]}"
return calls, remainder


def parse_tools(llm_response: str) -> list[tuple[str, Mapping]]:
"""A simple parser that will scan a string for tools and attempt to extract them; only works for json based outputs.
"""A simple parser that will scan a string for tools and attempt to extract them.

Recognizes JSON calls (`{"name": ..., "arguments": {...}}`, bare or in
`<tool_call>` tags) and Granite 4.2's XML format
(`<tool_call><function=NAME><parameter=KEY>VALUE</parameter></function></tool_call>`),
whose values are returned as strings. Markup quoted inside a call's
arguments is not parsed as a call.

Args:
llm_response: Raw string output from a language model.

Returns:
List of `(tool_name, arguments)` tuples for each tool call found.
List of `(tool_name, arguments)` tuples for each tool call found, XML
calls first.
"""
processed = " ".join(llm_response.split())
calls, _ = _parse_tool_calls(llm_response)
return [(name, args) for name, args, _ in calls]


def _parse_tool_calls(llm_response: str) -> tuple[list[tuple[str, Mapping, bool]], int]:
"""Like `parse_tools`, but also flags XML calls and counts unparsed `<tool_call>` tags.

Returns:
`(tool_name, arguments, from_xml)` per call, XML first, and the number
of `<tool_call>` tags that produced no call.
"""
# XML calls are blanked out of `remainder`, so their values aren't re-parsed as JSON.
xml_calls, remainder = _parse_xml_tool_calls(llm_response)
tools = [(name, args, True) for name, args in xml_calls]
processed = " ".join(remainder.split())

tools = []
for possible_tool in json_extraction(processed):
tool_name, tool_arguments = find_func(possible_tool)
if tool_name is not None and tool_arguments is not None:
tools.append((tool_name, tool_arguments))
return tools
tools.append((tool_name, tool_arguments, False))
return tools, _count_unparsed_tool_call_tags(remainder)


def _count_unparsed_tool_call_tags(text: str) -> int:
"""Count `<tool_call>` tags that open no JSON call, in text with XML calls blanked out.

Tags quoted in a JSON call's arguments, or wrapping a JSON call or list, don't count.
"""
if "<tool_call>" not in text:
return 0
json_spans = _json_tool_call_spans(text)
json_starts = {start for start, _ in json_spans}
unparsed = 0
for tag in re.finditer("<tool_call>", text):
if any(start <= tag.start() < end for start, end in json_spans):
continue
body = len(text) - len(text[tag.end() :].lstrip())
if body in json_starts or text.startswith("[", body):
continue
unparsed += 1
return unparsed


def validate_tool_arguments(
Expand All @@ -508,7 +613,8 @@ def validate_tool_arguments(
args: Raw arguments from model (post-JSON parsing)
coerce_types: If True, attempt type coercion for common cases (default: True)
strict: If True, raise ValidationError on failures; if False, log warnings
and return original args (default: False)
and return the failing args unvalidated, the rest validated
(default: False)

Returns:
Validated and optionally coerced arguments dict
Expand Down Expand Up @@ -791,11 +897,20 @@ def _build_pydantic_type_from_schema(schema: dict[str, Any]) -> Any:
MelleaLogger.get_logger().error(error_msg)
raise
else:
# Log warning and return original args
# Pass the failing arguments through as given; validate the rest.
MelleaLogger.get_logger().warning(
error_msg + "\nReturning original arguments without validation."
error_msg
+ "\nReturning those arguments as given; the rest were validated."
)
return dict(args)
failed = {error["loc"][0] for error in e.errors() if error["loc"]}
relaxed: dict[str, Any] = {
name: (Any, None) if name in failed else definition
for name, definition in field_definitions.items()
}
# Extra arguments are dumped as given (`extra="allow"`).
return create_model(
f"{tool_name}_Validator", __config__=model_config, **relaxed
)(**args).model_dump(exclude_unset=True)

except Exception as e:
# Catch any other errors during validation
Expand Down
67 changes: 62 additions & 5 deletions mellea/backends/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,15 +13,16 @@
from __future__ import annotations

import inspect
from collections.abc import Callable
import json
from collections.abc import Callable, Mapping, Sequence
from typing import Any

from ..core import Context, MelleaLogger, ModelToolCall, Span
from ..core.base import AbstractMelleaTool, ModelOutputThunk
from ..formatters import ChatFormatter
from ..helpers import merge_provider_fields
from ..stdlib.components import Message
from .tools import parse_tools, validate_tool_arguments
from .tools import _parse_tool_calls, validate_tool_arguments

# Chat = dict[Literal["role", "content"], str] # external apply_chat_template type hint is weaker
# Chat = dict[str, str | list[dict[str, Any]] ] # for multi-modal models
Expand Down Expand Up @@ -144,6 +145,50 @@ def to_chat(
return ctx_as_conversation


def _decode_text_args(
args: Mapping[str, Any], properties: Mapping[str, Any], required: Sequence[str]
) -> dict[str, Any]:
"""Decode XML-call values that stand for a list, dict or None.

The template writes lists and dicts as JSON and `None` as `None`. `None` or
`null` becomes None only for an optional parameter that can take it; other
scalars are left to `validate_tool_arguments`.

Args:
args: The call's arguments, as parsed from the model output.
properties: The tool schema's `properties`.
required: The tool schema's `required` parameter names.

Returns:
A new dict with the decoded values; `args` is not modified.
"""
decoded = dict(args)
for name, value in args.items():
if not isinstance(value, str):
continue
schema = properties.get(name) or {}
# `type` may be a list, at the top level or in an `anyOf` branch.
types: set[Any] = set()
for branch in [schema, *schema.get("anyOf", [])]:
declared = branch.get("type")
types.update(declared if isinstance(declared, list) else [declared])
if value.strip() in ("None", "null"):
# Nullable, or optional with no non-None default (a None default isn't emitted).
if name not in required and ("null" in types or "default" not in schema):
decoded[name] = None
continue
Comment thread
planetf1 marked this conversation as resolved.
if not types & {"object", "array"}:
continue
try:
parsed = json.loads(value)
except json.JSONDecodeError:
# Left for validate_tool_arguments to report.
continue
if isinstance(parsed, dict | list):
decoded[name] = parsed
return decoded


def to_tool_calls(
tools: dict[str, AbstractMelleaTool], decoded_result: str
) -> list[ModelToolCall] | None:
Expand All @@ -154,10 +199,17 @@ def to_tool_calls(
decoded_result: Raw model output string that may contain tool call markup.

Returns:
List of validated `ModelToolCall` (order preserved), or `None` if no tool calls were found.
List of validated `ModelToolCall` in the order parsed (XML-format calls
first, then JSON), or `None` if no tool calls were found.
"""
model_tool_calls: list[ModelToolCall] = []
for tool_name, tool_args in parse_tools(decoded_result):
parsed, unparsed = _parse_tool_calls(decoded_result)
if unparsed:
MelleaLogger.get_logger().warning(
f"model output contains {unparsed} <tool_call> block(s) that could "
"not be parsed as a tool call"
)
for tool_name, tool_args, from_xml in parsed:
func = tools.get(tool_name)
if func is None:
MelleaLogger.get_logger().warning(
Expand All @@ -167,9 +219,14 @@ def to_tool_calls(

# Clean up the function args slightly. Some models seem to
# hallucinate parameters when none are required.
param_map = func.as_json_tool["function"]["parameters"]["properties"]
parameters = func.as_json_tool["function"]["parameters"]
param_map = parameters["properties"]
if len(param_map) == 0:
tool_args = {}
elif from_xml: # only XML values arrive as text
tool_args = _decode_text_args(
tool_args, param_map, parameters.get("required") or []
)
Comment thread
planetf1 marked this conversation as resolved.

# Validate and coerce argument types
validated_args = validate_tool_arguments(func, tool_args, strict=False)
Expand Down
30 changes: 30 additions & 0 deletions test/backends/_granite_tokenizer.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,30 @@
# Copyright IBM Corp. All Rights Reserved.
# SPDX-License-Identifier: Apache-2.0

"""Load a cached Granite tokenizer, so tests can render the real chat template without torch or model weights."""

import pytest

_GRANITE_THINKING_MODEL_ID = "ibm-granite/granite-4.2-3b"


def _try_load_granite_tokenizer(model_id: str):
"""Return a Granite tokenizer from the local cache, or None.

A cached `config.json` does not guarantee the tokenizer's own files
(`tokenizer.json`, etc.) are also cached — e.g. after an interrupted or
partial download. `local_files_only=True` raises `OSError` in that case;
treat it the same as an absent cache rather than letting the test error.
Skips the calling test if `transformers` (the `mellea[hf]` extra) is missing.
"""
pytest.importorskip("transformers", reason="transformers not installed")
from huggingface_hub import _CACHED_NO_EXIST, try_to_load_from_cache
from transformers import AutoTokenizer

cached_config = try_to_load_from_cache(model_id, "config.json")
if cached_config is None or cached_config is _CACHED_NO_EXIST:
return None
try:
return AutoTokenizer.from_pretrained(model_id, local_files_only=True)
except OSError:
return None
25 changes: 4 additions & 21 deletions test/backends/test_huggingface_filter_options.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,10 @@
_HF_INTERNAL_TEMPLATE_VARS,
LocalHFBackend,
)
from test.backends._granite_tokenizer import (
_GRANITE_THINKING_MODEL_ID,
_try_load_granite_tokenizer,
)


def _make_backend(template: object) -> LocalHFBackend:
Expand Down Expand Up @@ -621,27 +625,6 @@ def test_generate_kwargs_allowlist_includes_known_generate_kwargs() -> None:
# ---------------------------------------------------------------------------

_GRANITE_MODEL_ID = "ibm-granite/granite-3.3-8b-instruct"
_GRANITE_THINKING_MODEL_ID = "ibm-granite/granite-4.2-3b"


def _try_load_granite_tokenizer(model_id: str):
"""Return a Granite tokenizer from the local cache, or None.

A cached `config.json` does not guarantee the tokenizer's own files
(`tokenizer.json`, etc.) are also cached — e.g. after an interrupted or
partial download. `local_files_only=True` raises `OSError` in that case;
treat it the same as an absent cache rather than letting the test error.
"""
from huggingface_hub import _CACHED_NO_EXIST, try_to_load_from_cache
from transformers import AutoTokenizer

cached_config = try_to_load_from_cache(model_id, "config.json")
if cached_config is None or cached_config is _CACHED_NO_EXIST:
return None
try:
return AutoTokenizer.from_pretrained(model_id, local_files_only=True)
except OSError:
return None


@pytest.mark.integration
Expand Down
Loading
Loading