From 833a5ad3eda7014ed2ae1fb8887520304220113b Mon Sep 17 00:00:00 2001 From: Nigel Jones Date: Tue, 29 Sep 2026 10:25:18 +0100 Subject: [PATCH 1/8] fix(backends): parse Granite XML tool calls on the local HF path LocalHFBackend parses tool calls with parse_tools(), which only understood JSON. From 4.2 the Granite chat template instructs the model to use an XML function format, so those calls were dropped silently and the markup reached the caller as response text. parse_tools() now extracts XML calls, returning values as raw strings, and blanks them out before the JSON pass so a JSON-looking parameter value is no longer parsed as a second call. to_tool_calls() JSON-decodes string values for object/array parameters and logs a warning when markup is present but nothing parses. Fixes #1689 Assisted-by: Claude Code Signed-off-by: Nigel Jones --- mellea/backends/tools.py | 40 +++++- mellea/backends/utils.py | 39 +++++- test/backends/test_huggingface_thinking.py | 41 ++++++ test/backends/test_tool_helpers.py | 67 +++++++++ test/backends/test_utils.py | 156 +++++++++++++++++++++ 5 files changed, 337 insertions(+), 6 deletions(-) diff --git a/mellea/backends/tools.py b/mellea/backends/tools.py index ae461dfb6a..7d71d6cfd2 100644 --- a/mellea/backends/tools.py +++ b/mellea/backends/tools.py @@ -471,18 +471,50 @@ def find_func(d: object) -> tuple[str | None, Mapping | None]: return None, None +# The XML function format Granite 4.2 chat templates instruct the model to use: +# \n\n\nVALUE\n\n\n +_XML_FUNCTION_RE = re.compile(r"\n]+)>(.*?)", re.DOTALL) +_XML_PARAMETER_RE = re.compile(r"\n]+)>(.*?)", re.DOTALL) + + +def _parse_xml_tool_calls(llm_response: str) -> list[tuple[str, Mapping]]: + """Extract tool calls written in the XML function format. + + Values are returned as raw strings, minus the single newline the template + wraps each one in; `to_tool_calls` coerces them against the tool schema. + """ + calls: list[tuple[str, Mapping]] = [] + for function in _XML_FUNCTION_RE.finditer(llm_response): + args = { + name.strip(): value.removeprefix("\n").removesuffix("\n") + for name, value in _XML_PARAMETER_RE.findall(function.group(2)) + } + calls.append((function.group(1).strip(), args)) + return calls + + 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. + + Two formats are recognized: JSON objects of the form + `{"name": ..., "arguments": {...}}` (bare, or inside `` tags), and + the XML function format that Granite 4.2 chat templates prescribe, + `VALUE` inside + `` tags. XML parameter values are returned as strings. 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()) + tools = _parse_xml_tool_calls(llm_response) + + # Blank out the XML calls so a JSON-looking parameter value is not parsed + # a second time as a call of its own. + processed = " ".join(_XML_FUNCTION_RE.sub(" ", llm_response).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: diff --git a/mellea/backends/utils.py b/mellea/backends/utils.py index d9fbaa0b96..7ef2205ef5 100644 --- a/mellea/backends/utils.py +++ b/mellea/backends/utils.py @@ -13,7 +13,8 @@ from __future__ import annotations import inspect -from collections.abc import Callable +import json +from collections.abc import Callable, Mapping from typing import Any from ..core import Context, MelleaLogger, ModelToolCall, Span @@ -144,6 +145,33 @@ def to_chat( return ctx_as_conversation +def _decode_json_container_args( + args: Mapping[str, Any], properties: Mapping[str, Any] +) -> dict[str, Any]: + """Decode string values of object/array parameters as JSON. + + Formats that carry every value as text (Granite's XML tool calls) render + list and dict arguments with `tojson`. Scalars are left for + `validate_tool_arguments` to coerce. + """ + decoded = dict(args) + for name, value in args.items(): + schema = properties.get(name) or {} + declared = schema.get("type") + types = set(declared) if isinstance(declared, list) else {declared} + types.update(s.get("type") for s in schema.get("anyOf", [])) + if not isinstance(value, str) or not types & {"object", "array"}: + continue + try: + parsed = json.loads(value) + except json.JSONDecodeError: + # Leave it as-is; validate_tool_arguments reports the mismatch. + 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: @@ -157,7 +185,12 @@ def to_tool_calls( List of validated `ModelToolCall` (order preserved), or `None` if no tool calls were found. """ model_tool_calls: list[ModelToolCall] = [] - for tool_name, tool_args in parse_tools(decoded_result): + parsed = parse_tools(decoded_result) + if not parsed and "" in decoded_result: + MelleaLogger.get_logger().warning( + "model output contains markup but no tool call could be parsed from it" + ) + for tool_name, tool_args in parsed: func = tools.get(tool_name) if func is None: MelleaLogger.get_logger().warning( @@ -170,6 +203,8 @@ def to_tool_calls( param_map = func.as_json_tool["function"]["parameters"]["properties"] if len(param_map) == 0: tool_args = {} + else: + tool_args = _decode_json_container_args(tool_args, param_map) # Validate and coerce argument types validated_args = validate_tool_arguments(func, tool_args, strict=False) diff --git a/test/backends/test_huggingface_thinking.py b/test/backends/test_huggingface_thinking.py index d602bc2b5d..bc05e69c8a 100644 --- a/test/backends/test_huggingface_thinking.py +++ b/test/backends/test_huggingface_thinking.py @@ -162,6 +162,47 @@ def fake_to_tool_calls(tools, text): assert recorded_text == ["call get_weather(city='Boston')"] +async def test_post_processing_registers_granite_xml_tool_call() -> None: + """A Granite 4.2 XML tool call after a thinking block becomes a ModelToolCall (#1689).""" + from mellea.backends.tools import MelleaTool + + def get_weather(city: str) -> str: + """Get the weather. + + Args: + city: the city name + """ + return f"sunny in {city}" + + backend = _make_backend() + mot = ModelOutputThunk( + value=( + "I should call the weather tool." + "\n\n\nBoston\n" + "\n\n" + ) + ) + mot._call.action = CBlock("What's the weather in Boston?") + mot._call.model_options = {} + + await backend.post_processing( + mot, + conversation=[], + _format=None, + tool_calls=True, + tools={"get_weather": MelleaTool.from_callable(get_weather)}, + seed=None, + input_ids=None, + ) + + assert mot.thinking == "I should call the weather tool." + assert mot.tool_calls is not None + assert [(c.name, c.args) for c in mot.tool_calls] == [ + ("get_weather", {"city": "Boston"}) + ] + assert mot.tool_calls[0].call_func() == "sunny in Boston" + + async def test_post_processing_skips_split_when_streaming() -> None: """Streaming generations must not have mot.value shrunk by post_processing. diff --git a/test/backends/test_tool_helpers.py b/test/backends/test_tool_helpers.py index a6c4bc611c..9eb4aad629 100644 --- a/test/backends/test_tool_helpers.py +++ b/test/backends/test_tool_helpers.py @@ -8,6 +8,7 @@ MelleaTool, add_tools_from_context_actions, add_tools_from_model_options, + parse_tools, ) from mellea.core import CBlock, Component, ModelOutputThunk, TemplateRepresentation @@ -109,5 +110,71 @@ def test_add_tools_from_context_actions(): assert tool2 == ftc1.tool2, f"{tool2} should == {ftc1.tool2}" +# --- parse_tools: Granite XML function format (#1689) --- + + +def test_parse_tools_granite_xml_single_call(): + """The shape the granite-4.2 chat template instructs the model to emit.""" + raw = ( + "\n\n\nBoston\n" + "\n\n" + ) + assert parse_tools(raw) == [("get_weather", {"city": "Boston"})] + + +def test_parse_tools_granite_xml_values_stay_raw_strings(): + """Values are returned as text; coercion against the schema happens later.""" + raw = ( + "\n\n\n3\n\n" + "\n[1, 2]\n\n\n" + ) + assert parse_tools(raw) == [("add", {"x": "3", "items": "[1, 2]"})] + + +def test_parse_tools_granite_xml_multiline_value_keeps_inner_newlines(): + """Only the newline the template wraps each value in is removed.""" + raw = ( + "\n\n\n" + "line one\n line two\n\n\n" + ) + assert parse_tools(raw) == [("write_file", {"body": "line one\n line two"})] + + +def test_parse_tools_granite_xml_multiple_calls_in_order(): + raw = ( + "\n\n\n1\n\n" + "\n\n" + "\n\n\n" + ) + assert parse_tools(raw) == [("first", {"a": "1"}), ("second", {})] + + +def test_parse_tools_granite_xml_after_reasoning_text(): + """The template allows natural-language reasoning before the call.""" + raw = ( + "I need the weather first.\n\n\n\n" + "\nBoston\n\n\n" + ) + assert parse_tools(raw) == [("get_weather", {"city": "Boston"})] + + +def test_parse_tools_granite_xml_json_value_is_not_a_second_call(): + """A parameter value that happens to look like a JSON tool call is data.""" + raw = ( + "\n\n\n" + '{"name": "add", "arguments": {"x": 1, "y": 2}}\n' + "\n\n" + ) + assert parse_tools(raw) == [ + ("log_event", {"payload": '{"name": "add", "arguments": {"x": 1, "y": 2}}'}) + ] + + +def test_parse_tools_json_inside_tool_call_tags_still_parses(): + """The Granite 4.0/4.1 shape: JSON inside `` tags.""" + raw = '\n{"name": "get_weather", "arguments": {"city": "Boston"}}\n' + assert parse_tools(raw) == [("get_weather", {"city": "Boston"})] + + if __name__ == "__main__": pytest.main([__file__]) diff --git a/test/backends/test_utils.py b/test/backends/test_utils.py index 4889d1b687..cbd968e4b1 100644 --- a/test/backends/test_utils.py +++ b/test/backends/test_utils.py @@ -3,6 +3,7 @@ """Unit tests for backends/utils.py — get_value accessor and to_tool_calls parser.""" +import logging from dataclasses import dataclass import pytest @@ -128,6 +129,161 @@ def test_to_tool_calls_string_arg_coerced_to_int(): assert result[0].args["y"] == 10 +# --- to_tool_calls: Granite XML function format (#1689) --- + + +def _tool_call_xml(name: str, args: dict[str, str]) -> str: + params = "".join(f"\n{v}\n\n" for k, v in args.items()) + return f"\n\n{params}\n" + + +def test_to_tool_calls_granite_xml_call_coerced_to_schema(): + registry = _make_tool_registry() + result = to_tool_calls(registry, _tool_call_xml("add", {"x": "3", "y": "4"})) + assert result is not None + assert len(result) == 1 + assert result[0].name == "add" + assert result[0].args == {"x": 3, "y": 4} + + +def test_to_tool_calls_granite_xml_array_param_decoded_from_json(): + """The template renders list/dict arguments with `tojson`, so they arrive as JSON text.""" + + def tag(labels: list[str]) -> str: + """Tag an item. + + Args: + labels: the labels to apply + """ + return ",".join(labels) + + registry = {"tag": MelleaTool.from_callable(tag)} + result = to_tool_calls(registry, _tool_call_xml("tag", {"labels": '["a", "b"]'})) + assert result is not None + assert result[0].args == {"labels": ["a", "b"]} + + +def test_to_tool_calls_json_text_for_string_param_stays_a_string(): + """Only object/array parameters are JSON-decoded.""" + + def note(text: str) -> str: + """Store a note. + + Args: + text: the note text + """ + return text + + registry = {"note": MelleaTool.from_callable(note)} + result = to_tool_calls(registry, _tool_call_xml("note", {"text": '{"a": 1}'})) + assert result is not None + assert result[0].args == {"text": '{"a": 1}'} + + +def test_to_tool_calls_granite_xml_optional_model_param_decoded(): + """`Optional[Model]` renders as `anyOf: [{type: object}, {type: null}]`.""" + from pydantic import BaseModel + + class Point(BaseModel): + x: int + y: int + + def plot(point: Point | None = None) -> str: + """Plot a point. + + Args: + point: the point to plot + """ + return str(point) + + registry = {"plot": MelleaTool.from_callable(plot)} + raw = _tool_call_xml("plot", {"point": '{"x": 1, "y": 2}'}) + result = to_tool_calls(registry, raw) + assert result is not None + assert result[0].args == {"point": {"x": 1, "y": 2}} + + +def test_to_tool_calls_undecodable_array_value_left_for_validation(): + """Text that is not JSON is passed through; validation reports the mismatch.""" + + def tag(labels: list[str]) -> str: + """Tag an item. + + Args: + labels: the labels to apply + """ + return ",".join(labels) + + registry = {"tag": MelleaTool.from_callable(tag)} + result = to_tool_calls(registry, _tool_call_xml("tag", {"labels": "a, b"})) + assert result is not None + assert result[0].args == {"labels": "a, b"} + + +def test_to_tool_calls_warns_on_unparsable_tool_call_markup( + caplog: pytest.LogCaptureFixture, +): + registry = _make_tool_registry() + with caplog.at_level(logging.WARNING, logger="mellea"): + result = to_tool_calls( + registry, "\nadd three and four\n" + ) + assert result is None + assert any("" in r.getMessage() for r in caplog.records) + + +@pytest.mark.integration +def test_to_tool_calls_round_trips_real_granite_template() -> None: + """Parse the tool-call markup the real granite-4.2-3b chat template renders. + + Loads the template only (no GPU, no model weights); skips if not locally cached. + """ + from test.backends.test_huggingface_filter_options import ( + _GRANITE_THINKING_MODEL_ID, + _try_load_granite_tokenizer, + ) + + tok = _try_load_granite_tokenizer(_GRANITE_THINKING_MODEL_ID) + if tok is None: + pytest.skip(f"{_GRANITE_THINKING_MODEL_ID} not in local HF cache") + + def tag(labels: list[str], note: str) -> str: + """Tag an item. + + Args: + labels: the labels to apply + note: a free-text note + """ + return note + + registry = {**_make_tool_registry(), "tag": MelleaTool.from_callable(tag)} + calls = [ + {"name": "add", "arguments": {"x": 3, "y": 4}}, + {"name": "tag", "arguments": {"labels": ["a", "b"], "note": "first\nsecond"}}, + ] + rendered = tok.apply_chat_template( + [ + {"role": "user", "content": "Add 3 and 4, then tag it."}, + { + "role": "assistant", + "content": "", + "tool_calls": [ + {"id": str(i), "type": "function", "function": c} + for i, c in enumerate(calls) + ], + }, + ], + tokenize=False, + ) + assert "" in rendered # the template really uses the XML format + + result = to_tool_calls(registry, rendered) + assert result is not None + assert [(r.name, r.args) for r in result] == [ + (c["name"], c["arguments"]) for c in calls + ] + + # --- to_chat --- From 46007cbc8124af753009a045175782d466e7e2d7 Mon Sep 17 00:00:00 2001 From: Nigel Jones Date: Tue, 29 Sep 2026 11:53:59 +0100 Subject: [PATCH 2/8] fix(backends): accept list-valued types inside anyOf when decoding tool args Schemas brought in through MelleaTool.from_langchain() or from_smolagents() can carry an anyOf branch whose type is a list, such as ["array", "null"]. Collecting those into a set raised TypeError and failed the whole generation. Normalise each branch's type the same way as the top-level one. Refs #1689 Assisted-by: Claude Code Signed-off-by: Nigel Jones --- mellea/backends/utils.py | 8 +++++--- test/backends/test_utils.py | 8 ++++++++ 2 files changed, 13 insertions(+), 3 deletions(-) diff --git a/mellea/backends/utils.py b/mellea/backends/utils.py index 7ef2205ef5..d70d2e45d4 100644 --- a/mellea/backends/utils.py +++ b/mellea/backends/utils.py @@ -157,9 +157,11 @@ def _decode_json_container_args( decoded = dict(args) for name, value in args.items(): schema = properties.get(name) or {} - declared = schema.get("type") - types = set(declared) if isinstance(declared, list) else {declared} - types.update(s.get("type") for s in schema.get("anyOf", [])) + # `type` may be a list (`["array", "null"]`), 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 not isinstance(value, str) or not types & {"object", "array"}: continue try: diff --git a/test/backends/test_utils.py b/test/backends/test_utils.py index cbd968e4b1..2b013ca5d7 100644 --- a/test/backends/test_utils.py +++ b/test/backends/test_utils.py @@ -203,6 +203,14 @@ def plot(point: Point | None = None) -> str: assert result[0].args == {"point": {"x": 1, "y": 2}} +def test_decode_json_container_args_list_type_inside_any_of(): + """External schemas (LangChain, smolagents) can put a list-valued `type` in an `anyOf` branch.""" + from mellea.backends.utils import _decode_json_container_args + + properties = {"x": {"anyOf": [{"type": ["array", "null"]}]}} + assert _decode_json_container_args({"x": "[1, 2]"}, properties) == {"x": [1, 2]} + + def test_to_tool_calls_undecodable_array_value_left_for_validation(): """Text that is not JSON is passed through; validation reports the mismatch.""" From 0f1abedd498ed0de568c9712b768332e505418f0 Mon Sep 17 00:00:00 2001 From: Nigel Jones Date: Tue, 29 Sep 2026 12:15:24 +0100 Subject: [PATCH 3/8] fix(backends): scope Granite XML tool-call parsing to real calls Code review found three ways the XML pass could produce the wrong call: - XML markup inside a JSON call's string argument was parsed as a second call and blanked out of the argument. JSON tool-call spans are now located first and XML matches inside them are skipped. - A call missing `` ran on into the next call and took its parameters. The body can no longer cross another `` tag. - A `` mentioned in text before the real call started a match there, so the wrong tool was called. Matches must now open with ``, as the template requires; the closing tag stays optional for truncated output. to_tool_calls() now warns whenever there are more `` blocks than parsed calls, so a dropped call is logged even when others parse. Refs #1689 Assisted-by: Claude Code Signed-off-by: Nigel Jones --- mellea/backends/tools.py | 75 +++++++++++++++++++++++++----- mellea/backends/utils.py | 6 ++- test/backends/test_tool_helpers.py | 47 +++++++++++++++++++ test/backends/test_utils.py | 16 +++++++ 4 files changed, 130 insertions(+), 14 deletions(-) diff --git a/mellea/backends/tools.py b/mellea/backends/tools.py index 7d71d6cfd2..75960e2928 100644 --- a/mellea/backends/tools.py +++ b/mellea/backends/tools.py @@ -473,24 +473,75 @@ def find_func(d: object) -> tuple[str | None, Mapping | None]: # The XML function format Granite 4.2 chat templates instruct the model to use: # \n\n\nVALUE\n\n\n -_XML_FUNCTION_RE = re.compile(r"\n]+)>(.*?)", re.DOTALL) +# A call must open with ``, and its body may not run into another call, so a +# call missing `` (or a `` is optional, for truncated output. +# The format has no escaping, so a value that itself contains `` or +# `` is cut short there, and one containing `` +# tag is not parsed at all (`to_tool_calls` logs the dropped call). +_XML_TOOL_CALL_RE = re.compile( + r"\s*\n]+)>" + r"((?:(?!).)*?)" + r"\s*(?:)?", + re.DOTALL, +) _XML_PARAMETER_RE = re.compile(r"\n]+)>(.*?)", re.DOTALL) -def _parse_xml_tool_calls(llm_response: str) -> list[tuple[str, Mapping]]: +def _json_tool_call_spans(text: str) -> list[tuple[int, int]]: + """Return the `(start, end)` offsets of the JSON tool calls in `text`. + + Accepts what `json_extraction` and `find_func` accept, but on the raw text, + so the offsets line up with it. + """ + spans: list[tuple[int, int]] = [] + 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 tool calls written in the XML function format. Values are returned as raw strings, minus the single newline the template wraps each one in; `to_tool_calls` coerces them against the tool schema. + Markup inside the arguments of a JSON tool call is data and is skipped. + + Returns: + The calls, and `llm_response` with those calls 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 function in _XML_FUNCTION_RE.finditer(llm_response): + for match in matches: args = { name.strip(): value.removeprefix("\n").removesuffix("\n") - for name, value in _XML_PARAMETER_RE.findall(function.group(2)) + for name, value in _XML_PARAMETER_RE.findall(match.group(2)) } - calls.append((function.group(1).strip(), args)) - return calls + 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]]: @@ -500,7 +551,8 @@ def parse_tools(llm_response: str) -> list[tuple[str, Mapping]]: `{"name": ..., "arguments": {...}}` (bare, or inside `` tags), and the XML function format that Granite 4.2 chat templates prescribe, `VALUE` inside - `` tags. XML parameter values are returned as strings. + `` tags. XML parameter values are returned as strings. XML + markup that appears inside a JSON call's arguments is treated as data. Args: llm_response: Raw string output from a language model. @@ -509,11 +561,10 @@ def parse_tools(llm_response: str) -> list[tuple[str, Mapping]]: List of `(tool_name, arguments)` tuples for each tool call found, XML calls first. """ - tools = _parse_xml_tool_calls(llm_response) - - # Blank out the XML calls so a JSON-looking parameter value is not parsed - # a second time as a call of its own. - processed = " ".join(_XML_FUNCTION_RE.sub(" ", llm_response).split()) + # The XML calls come back blanked out of `remainder`, so a JSON-looking + # parameter value is not parsed a second time as a call of its own. + tools, remainder = _parse_xml_tool_calls(llm_response) + processed = " ".join(remainder.split()) for possible_tool in json_extraction(processed): tool_name, tool_arguments = find_func(possible_tool) diff --git a/mellea/backends/utils.py b/mellea/backends/utils.py index d70d2e45d4..8bf4f9155a 100644 --- a/mellea/backends/utils.py +++ b/mellea/backends/utils.py @@ -188,9 +188,11 @@ def to_tool_calls( """ model_tool_calls: list[ModelToolCall] = [] parsed = parse_tools(decoded_result) - if not parsed and "" in decoded_result: + tagged = decoded_result.count("") + if tagged > len(parsed): MelleaLogger.get_logger().warning( - "model output contains markup but no tool call could be parsed from it" + f"model output contains {tagged} block(s) but only " + f"{len(parsed)} tool call(s) could be parsed from it" ) for tool_name, tool_args in parsed: func = tools.get(tool_name) diff --git a/test/backends/test_tool_helpers.py b/test/backends/test_tool_helpers.py index 9eb4aad629..7b31f56558 100644 --- a/test/backends/test_tool_helpers.py +++ b/test/backends/test_tool_helpers.py @@ -170,6 +170,53 @@ def test_parse_tools_granite_xml_json_value_is_not_a_second_call(): ] +_XML_MARKUP_AS_DATA = ( + "Granite writes /" + " for a call" +) + + +@pytest.mark.parametrize("wrapped", [False, True]) +def test_parse_tools_json_argument_containing_xml_markup_is_data(wrapped: bool): + """XML call markup inside a JSON string argument is data, not a second call.""" + import json + + call = json.dumps({"name": "write_doc", "arguments": {"text": _XML_MARKUP_AS_DATA}}) + raw = f"\n{call}\n" if wrapped else call + assert parse_tools(raw) == [("write_doc", {"text": _XML_MARKUP_AS_DATA})] + + +def test_parse_tools_granite_xml_missing_close_does_not_merge_calls(): + """A call missing `` must not swallow the next call's parameters.""" + raw = ( + "foo" + "Boston" + "" + ) + assert parse_tools(raw) == [("get_weather", {"city": "Boston"})] + + +def test_parse_tools_granite_xml_prose_mention_is_not_a_call(): + """`` mentioned in text before the call is not a call of its own.""" + raw = ( + "Maybe is wrong.\n\n\n" + "\nBoston\n\n\n" + ) + assert parse_tools(raw) == [("get_weather", {"city": "Boston"})] + + +def test_parse_tools_granite_xml_without_closing_tool_call_tag(): + """Output that stops after `` still parses.""" + raw = "\n\n\nBoston\n\n" + assert parse_tools(raw) == [("get_weather", {"city": "Boston"})] + + +def test_parse_tools_xml_function_block_needs_tool_call_wrapper(): + """The template nests every call in `` tags; a bare block is not a call.""" + raw = "\n\nBoston\n\n" + assert parse_tools(raw) == [] + + def test_parse_tools_json_inside_tool_call_tags_still_parses(): """The Granite 4.0/4.1 shape: JSON inside `` tags.""" raw = '\n{"name": "get_weather", "arguments": {"city": "Boston"}}\n' diff --git a/test/backends/test_utils.py b/test/backends/test_utils.py index 2b013ca5d7..9ff2400010 100644 --- a/test/backends/test_utils.py +++ b/test/backends/test_utils.py @@ -240,6 +240,22 @@ def test_to_tool_calls_warns_on_unparsable_tool_call_markup( assert any("" in r.getMessage() for r in caplog.records) +def test_to_tool_calls_warns_when_some_tool_calls_are_dropped( + caplog: pytest.LogCaptureFixture, +): + """A malformed call next to a good one is dropped, and the drop is logged.""" + registry = _make_tool_registry() + raw = ( + "1" + + _tool_call_xml("greet", {"name": "Ada"}) + ) + with caplog.at_level(logging.WARNING, logger="mellea"): + result = to_tool_calls(registry, raw) + assert result is not None + assert [(r.name, r.args) for r in result] == [("greet", {"name": "Ada"})] + assert any("" in r.getMessage() for r in caplog.records) + + @pytest.mark.integration def test_to_tool_calls_round_trips_real_granite_template() -> None: """Parse the tool-call markup the real granite-4.2-3b chat template renders. From 25d82fe5494af55e39ab4b49d24eae0b56866f66 Mon Sep 17 00:00:00 2001 From: Nigel Jones Date: Tue, 29 Sep 2026 12:38:21 +0100 Subject: [PATCH 4/8] fix(backends): read None/null text as null for optional tool arguments Granite's chat template renders a Python None as the text "None", and XML tool calls carry every value as text. So "None" or "null" for an optional parameter now becomes None before validation, instead of reaching the tool as a string. Required parameters are left as sent. Validation of an explicit None for an optional parameter is fixed separately in the tool-schema work for #1693. Refs #1689 Assisted-by: Claude Code Signed-off-by: Nigel Jones --- mellea/backends/utils.py | 28 ++++++++++++++++++-------- test/backends/test_utils.py | 39 ++++++++++++++++++++++++++++++++++--- 2 files changed, 56 insertions(+), 11 deletions(-) diff --git a/mellea/backends/utils.py b/mellea/backends/utils.py index 8bf4f9155a..82a7c38499 100644 --- a/mellea/backends/utils.py +++ b/mellea/backends/utils.py @@ -14,7 +14,7 @@ import inspect import json -from collections.abc import Callable, Mapping +from collections.abc import Callable, Mapping, Sequence from typing import Any from ..core import Context, MelleaLogger, ModelToolCall, Span @@ -145,17 +145,26 @@ def to_chat( return ctx_as_conversation -def _decode_json_container_args( - args: Mapping[str, Any], properties: Mapping[str, Any] +def _decode_text_args( + args: Mapping[str, Any], properties: Mapping[str, Any], required: Sequence[str] ) -> dict[str, Any]: - """Decode string values of object/array parameters as JSON. + """Decode argument values that arrive as text but stand for something else. Formats that carry every value as text (Granite's XML tool calls) render - list and dict arguments with `tojson`. Scalars are left for - `validate_tool_arguments` to coerce. + list and dict arguments with `tojson` and a Python `None` as `None`. So a + string value for an object/array parameter is decoded as JSON, and `None` + or `null` for an optional parameter becomes `None`. Other scalars are left + for `validate_tool_arguments` to coerce. """ decoded = dict(args) for name, value in args.items(): + if ( + name not in required + and isinstance(value, str) + and value.strip() in ("None", "null") + ): + decoded[name] = None + continue schema = properties.get(name) or {} # `type` may be a list (`["array", "null"]`), at the top level or in an `anyOf` branch. types: set[Any] = set() @@ -204,11 +213,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 = {} else: - tool_args = _decode_json_container_args(tool_args, param_map) + tool_args = _decode_text_args( + tool_args, param_map, parameters.get("required") or [] + ) # Validate and coerce argument types validated_args = validate_tool_arguments(func, tool_args, strict=False) diff --git a/test/backends/test_utils.py b/test/backends/test_utils.py index 9ff2400010..52deaa5a1f 100644 --- a/test/backends/test_utils.py +++ b/test/backends/test_utils.py @@ -203,12 +203,45 @@ def plot(point: Point | None = None) -> str: assert result[0].args == {"point": {"x": 1, "y": 2}} -def test_decode_json_container_args_list_type_inside_any_of(): +def test_decode_text_args_list_type_inside_any_of(): """External schemas (LangChain, smolagents) can put a list-valued `type` in an `anyOf` branch.""" - from mellea.backends.utils import _decode_json_container_args + from mellea.backends.utils import _decode_text_args properties = {"x": {"anyOf": [{"type": ["array", "null"]}]}} - assert _decode_json_container_args({"x": "[1, 2]"}, properties) == {"x": [1, 2]} + assert _decode_text_args({"x": "[1, 2]"}, properties, required=[]) == {"x": [1, 2]} + + +@pytest.mark.parametrize("text", ["None", "null", " None "]) +def test_to_tool_calls_granite_xml_none_text_for_optional_param_is_null(text: str): + """The template renders a Python `None` as `None`; for an optional parameter that is a null.""" + + def search(query: str, limit: int | None = None) -> str: + """Search. + + Args: + query: what to search for + limit: optional result limit + """ + return query + + from mellea.core.base import AbstractMelleaTool + + registry: dict[str, AbstractMelleaTool] = { + "search": MelleaTool.from_callable(search) + } + raw = _tool_call_xml("search", {"query": "cats", "limit": text}) + result = to_tool_calls(registry, raw) + assert result is not None + assert result[0].args["query"] == "cats" + assert result[0].args["limit"] is None + + +def test_to_tool_calls_none_text_for_required_param_is_kept(): + """A required parameter has no null to map to, so the text stays as sent.""" + registry = _make_tool_registry() + result = to_tool_calls(registry, _tool_call_xml("greet", {"name": "None"})) + assert result is not None + assert result[0].args == {"name": "None"} def test_to_tool_calls_undecodable_array_value_left_for_validation(): From 80efe417a9795f661236972cdade42188dd9bf66 Mon Sep 17 00:00:00 2001 From: Nigel Jones Date: Tue, 29 Sep 2026 12:38:48 +0100 Subject: [PATCH 5/8] test: share the Granite tokenizer loader between backend tests _try_load_granite_tokenizer lived in test_huggingface_filter_options.py, which skips without torch and imports the HF backend at module level. test_utils.py and test_huggingface_thinking.py imported it from there. Move it, with the model id, into test/backends/_granite_tokenizer.py, which needs only transformers and skips cleanly without it. Refs #1689 Assisted-by: Claude Code Signed-off-by: Nigel Jones --- test/backends/_granite_tokenizer.py | 35 +++++++++++++++++++ .../test_huggingface_filter_options.py | 25 +++---------- test/backends/test_huggingface_thinking.py | 2 +- test/backends/test_utils.py | 9 +++-- 4 files changed, 44 insertions(+), 27 deletions(-) create mode 100644 test/backends/_granite_tokenizer.py diff --git a/test/backends/_granite_tokenizer.py b/test/backends/_granite_tokenizer.py new file mode 100644 index 0000000000..1104ddfb15 --- /dev/null +++ b/test/backends/_granite_tokenizer.py @@ -0,0 +1,35 @@ +# Copyright IBM Corp. All Rights Reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Load a cached Granite tokenizer, for tests that render the real chat template. + +Kept apart from the HF backend tests so that a light test file can use the +template without importing torch or `mellea.backends.huggingface`. Only the +tokenizer files are loaded; no model weights and no GPU. +""" + +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 diff --git a/test/backends/test_huggingface_filter_options.py b/test/backends/test_huggingface_filter_options.py index 6aecdf010b..f6af02561b 100644 --- a/test/backends/test_huggingface_filter_options.py +++ b/test/backends/test_huggingface_filter_options.py @@ -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: @@ -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 diff --git a/test/backends/test_huggingface_thinking.py b/test/backends/test_huggingface_thinking.py index bc05e69c8a..47a7861a42 100644 --- a/test/backends/test_huggingface_thinking.py +++ b/test/backends/test_huggingface_thinking.py @@ -17,7 +17,7 @@ from mellea.backends import ModelOption from mellea.backends.huggingface import LocalHFBackend, _split_think_tags from mellea.core.base import CBlock, ModelOutputThunk -from test.backends.test_huggingface_filter_options import ( +from test.backends._granite_tokenizer import ( _GRANITE_THINKING_MODEL_ID, _try_load_granite_tokenizer, ) diff --git a/test/backends/test_utils.py b/test/backends/test_utils.py index 52deaa5a1f..9439a49e38 100644 --- a/test/backends/test_utils.py +++ b/test/backends/test_utils.py @@ -16,6 +16,10 @@ ) from mellea.core import ModelToolCall from mellea.core.base import GenerationMetadata, ModelOutputThunk +from test.backends._granite_tokenizer import ( + _GRANITE_THINKING_MODEL_ID, + _try_load_granite_tokenizer, +) # --- get_value --- @@ -295,11 +299,6 @@ def test_to_tool_calls_round_trips_real_granite_template() -> None: Loads the template only (no GPU, no model weights); skips if not locally cached. """ - from test.backends.test_huggingface_filter_options import ( - _GRANITE_THINKING_MODEL_ID, - _try_load_granite_tokenizer, - ) - tok = _try_load_granite_tokenizer(_GRANITE_THINKING_MODEL_ID) if tok is None: pytest.skip(f"{_GRANITE_THINKING_MODEL_ID} not in local HF cache") From 5ea0374498f7726552b6df333ae1438aaf32e9e7 Mon Sep 17 00:00:00 2001 From: Nigel Jones Date: Tue, 29 Sep 2026 13:38:49 +0100 Subject: [PATCH 6/8] fix(backends): tighten Granite XML tool-call parsing after a second review A second three-reviewer round found: - to_tool_calls() text-decoded every parsed call, so a JSON call with an optional string argument of "None" reached the tool as None. Decoding now applies only to XML-format calls, which are the ones that carry every value as text. "None"/"null" also becomes None only where the parameter can take it: optional, and either nullable in the schema or with no default other than None. `title: str = "Untitled"` keeps the text. - An XML value quoting a whole `{json}` made the outer call fail to match, and the quoted JSON then ran as a call. A call body must now be parameter blocks separated by whitespace, so a value can hold any markup except ``. - The dropped-call warning compared a raw tag count with the number of parsed calls, so it fired on tags quoted in a JSON argument and missed a drop hidden by another call. It now counts `` tags that open no parsed call. Refs #1689 Assisted-by: Claude Code Signed-off-by: Nigel Jones --- mellea/backends/tools.py | 58 +++++++++++++++++---- mellea/backends/utils.py | 46 ++++++++++------- test/backends/test_tool_helpers.py | 20 ++++++++ test/backends/test_utils.py | 81 ++++++++++++++++++++++++++++++ 4 files changed, 176 insertions(+), 29 deletions(-) diff --git a/mellea/backends/tools.py b/mellea/backends/tools.py index 75960e2928..0bfe50fd16 100644 --- a/mellea/backends/tools.py +++ b/mellea/backends/tools.py @@ -473,15 +473,15 @@ def find_func(d: object) -> tuple[str | None, Mapping | None]: # The XML function format Granite 4.2 chat templates instruct the model to use: # \n\n\nVALUE\n\n\n -# A call must open with ``, and its body may not run into another call, so a -# call missing `` (or a `` is optional, for truncated output. -# The format has no escaping, so a value that itself contains `` or -# `` is cut short there, and one containing `` -# tag is not parsed at all (`to_tool_calls` logs the dropped call). +# A call must open with ``, and its body must be parameter blocks separated +# only by whitespace. So a value can hold any text, markup included, except +# ``; a call missing `` can't run into the next call; and a +# `` is optional, +# for truncated output. The format has no escaping, so a value that contains +# `` leaves its call unparsed (`to_tool_calls` logs the dropped call). _XML_TOOL_CALL_RE = re.compile( - r"\s*\n]+)>" - r"((?:(?!).)*?)" + r"\s*\n]+)>\s*" + r"((?:\n]+>(?:(?!).)*\s*)*)" r"\s*(?:)?", re.DOTALL, ) @@ -495,6 +495,8 @@ def _json_tool_call_spans(text: str) -> list[tuple[int, int]]: so the offsets line up with it. """ spans: list[tuple[int, int]] = [] + # `json_extraction` runs after whitespace is collapsed; this runs on the raw + # text, so it has to accept literal newlines inside strings. decoder = json.JSONDecoder(strict=False) index = text.find("{") while index != -1: @@ -561,16 +563,50 @@ def parse_tools(llm_response: str) -> list[tuple[str, Mapping]]: List of `(tool_name, arguments)` tuples for each tool call found, XML calls first. """ + 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]: + """Parse tool calls as `parse_tools` does, also saying which came from XML. + + Returns: + The calls as `(tool_name, arguments, from_xml)`, XML calls first, and + the number of `` tags that open none of them. + """ # The XML calls come back blanked out of `remainder`, so a JSON-looking # parameter value is not parsed a second time as a call of its own. - tools, remainder = _parse_xml_tool_calls(llm_response) + xml_calls, remainder = _parse_xml_tool_calls(llm_response) + tools = [(name, args, True) for name, args in xml_calls] processed = " ".join(remainder.split()) 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 the `` tags in `text` that don't open a JSON tool call. + + Run on the output with the parsed XML calls blanked out. A tag quoted inside + a JSON call's arguments is data, and one directly followed by a JSON call + (or a JSON list of calls) is that call's wrapper; neither is counted. + """ + if "" 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("", 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( diff --git a/mellea/backends/utils.py b/mellea/backends/utils.py index 82a7c38499..f62c944d40 100644 --- a/mellea/backends/utils.py +++ b/mellea/backends/utils.py @@ -22,7 +22,7 @@ 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 @@ -153,17 +153,20 @@ def _decode_text_args( Formats that carry every value as text (Granite's XML tool calls) render list and dict arguments with `tojson` and a Python `None` as `None`. So a string value for an object/array parameter is decoded as JSON, and `None` - or `null` for an optional parameter becomes `None`. Other scalars are left - for `validate_tool_arguments` to coerce. + or `null` becomes `None` for an optional parameter that can take it. Other + scalars are left for `validate_tool_arguments` to coerce. + + 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 ( - name not in required - and isinstance(value, str) - and value.strip() in ("None", "null") - ): - decoded[name] = None + if not isinstance(value, str): continue schema = properties.get(name) or {} # `type` may be a list (`["array", "null"]`), at the top level or in an `anyOf` branch. @@ -171,7 +174,14 @@ def _decode_text_args( for branch in [schema, *schema.get("anyOf", [])]: declared = branch.get("type") types.update(declared if isinstance(declared, list) else [declared]) - if not isinstance(value, str) or not types & {"object", "array"}: + if value.strip() in ("None", "null"): + # Only where None is a valid value: the schema allows null, or the + # parameter is optional with no default other than None (the emitted + # schema leaves a None default out). + if name not in required and ("null" in types or "default" not in schema): + decoded[name] = None + continue + if not types & {"object", "array"}: continue try: parsed = json.loads(value) @@ -193,17 +203,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] = [] - parsed = parse_tools(decoded_result) - tagged = decoded_result.count("") - if tagged > len(parsed): + parsed, unparsed = _parse_tool_calls(decoded_result) + if unparsed: MelleaLogger.get_logger().warning( - f"model output contains {tagged} block(s) but only " - f"{len(parsed)} tool call(s) could be parsed from it" + f"model output contains {unparsed} block(s) that could " + "not be parsed as a tool call" ) - for tool_name, tool_args in parsed: + for tool_name, tool_args, from_xml in parsed: func = tools.get(tool_name) if func is None: MelleaLogger.get_logger().warning( @@ -217,7 +227,7 @@ def to_tool_calls( param_map = parameters["properties"] if len(param_map) == 0: tool_args = {} - else: + elif from_xml: # XML calls carry every value as text; JSON calls are typed tool_args = _decode_text_args( tool_args, param_map, parameters.get("required") or [] ) diff --git a/test/backends/test_tool_helpers.py b/test/backends/test_tool_helpers.py index 7b31f56558..0cf895ac2d 100644 --- a/test/backends/test_tool_helpers.py +++ b/test/backends/test_tool_helpers.py @@ -205,6 +205,26 @@ def test_parse_tools_granite_xml_prose_mention_is_not_a_call(): assert parse_tools(raw) == [("get_weather", {"city": "Boston"})] +def test_parse_tools_granite_xml_quoted_tool_call_in_value_is_data(): + """A value quoting a whole JSON tool call, tags included, is not run as a call.""" + quoted = '{"name": "add", "arguments": {"x": 1, "y": 2}}' + raw = ( + "\n\n\n" + f"{quoted}\n\n\n" + ) + assert parse_tools(raw) == [("log_event", {"payload": quoted})] + + +def test_parse_tools_granite_xml_value_may_contain_function_markup(): + """`` and `` inside a value are kept as text.""" + text = "call and close it with " + raw = ( + "\n\n\n" + f"{text}\n\n\n" + ) + assert parse_tools(raw) == [("write_doc", {"text": text})] + + def test_parse_tools_granite_xml_without_closing_tool_call_tag(): """Output that stops after `` still parses.""" raw = "\n\n\nBoston\n\n" diff --git a/test/backends/test_utils.py b/test/backends/test_utils.py index 9439a49e38..fd1fcc0283 100644 --- a/test/backends/test_utils.py +++ b/test/backends/test_utils.py @@ -248,6 +248,23 @@ def test_to_tool_calls_none_text_for_required_param_is_kept(): assert result[0].args == {"name": "None"} +def test_to_tool_calls_none_text_kept_when_param_has_a_real_default(): + """`title: str = "Untitled"` can't be None, so "None" stays text rather than overriding the default.""" + + def make(title: str = "Untitled") -> str: + """Make something. + + Args: + title: its title + """ + return title + + registry = {"make": MelleaTool.from_callable(make)} + result = to_tool_calls(registry, _tool_call_xml("make", {"title": "None"})) + assert result is not None + assert result[0].args == {"title": "None"} + + def test_to_tool_calls_undecodable_array_value_left_for_validation(): """Text that is not JSON is passed through; validation reports the mismatch.""" @@ -277,6 +294,70 @@ def test_to_tool_calls_warns_on_unparsable_tool_call_markup( assert any("" in r.getMessage() for r in caplog.records) +def test_to_tool_calls_no_warning_for_tool_call_text_inside_a_json_argument( + caplog: pytest.LogCaptureFixture, +): + """`` quoted inside a parsed call's argument is data, not a dropped call.""" + import json + + def write_doc(text: str) -> str: + """Write a document. + + Args: + text: the document text + """ + return text + + from mellea.core.base import AbstractMelleaTool + + registry: dict[str, AbstractMelleaTool] = { + "write_doc": MelleaTool.from_callable(write_doc) + } + text = "wrap each call in tags" + call = json.dumps({"name": "write_doc", "arguments": {"text": text}}) + with caplog.at_level(logging.WARNING, logger="mellea"): + result = to_tool_calls(registry, f"\n{call}\n") + assert result is not None + assert [(r.name, r.args) for r in result] == [("write_doc", {"text": text})] + assert not caplog.records + + +def test_to_tool_calls_json_call_arguments_are_not_text_decoded(): + """Only XML calls carry values as text; a JSON call's "None" string stays a string.""" + import json + + def search(query: str, note: str | None = None) -> str: + """Search. + + Args: + query: what to search for + note: an optional note + """ + return query + + registry = {"search": MelleaTool.from_callable(search)} + raw = json.dumps({"name": "search", "arguments": {"query": "cats", "note": "None"}}) + result = to_tool_calls(registry, raw) + assert result is not None + assert result[0].args == {"query": "cats", "note": "None"} + + +def test_to_tool_calls_warns_when_a_bare_json_call_hides_a_dropped_one( + caplog: pytest.LogCaptureFixture, +): + """A dropped tagged call is logged even when another call keeps the counts level.""" + import json + + registry = _make_tool_registry() + good = json.dumps({"name": "greet", "arguments": {"name": "Ada"}}) + raw = f"\ngreet Ada\n\n{good}" + with caplog.at_level(logging.WARNING, logger="mellea"): + result = to_tool_calls(registry, raw) + assert result is not None + assert [(r.name, r.args) for r in result] == [("greet", {"name": "Ada"})] + assert any("" in r.getMessage() for r in caplog.records) + + def test_to_tool_calls_warns_when_some_tool_calls_are_dropped( caplog: pytest.LogCaptureFixture, ): From f3bda2456be5395cad65519ea26977310ba74edd Mon Sep 17 00:00:00 2001 From: Nigel Jones Date: Tue, 29 Sep 2026 13:47:43 +0100 Subject: [PATCH 7/8] docs(backends): tighten comments in the Granite XML tool-call parsing Comment and docstring changes only; the code is unchanged. Refs #1689 Assisted-by: Claude Code Signed-off-by: Nigel Jones --- mellea/backends/tools.py | 54 +++++++++++------------------ mellea/backends/utils.py | 20 +++++------ test/backends/_granite_tokenizer.py | 7 +--- 3 files changed, 29 insertions(+), 52 deletions(-) diff --git a/mellea/backends/tools.py b/mellea/backends/tools.py index 0bfe50fd16..e1884be529 100644 --- a/mellea/backends/tools.py +++ b/mellea/backends/tools.py @@ -471,14 +471,11 @@ def find_func(d: object) -> tuple[str | None, Mapping | None]: return None, None -# The XML function format Granite 4.2 chat templates instruct the model to use: +# Granite 4.2's XML tool-call format: # \n\n\nVALUE\n\n\n -# A call must open with ``, and its body must be parameter blocks separated -# only by whitespace. So a value can hold any text, markup included, except -# ``; a call missing `` can't run into the next call; and a -# `` is optional, -# for truncated output. The format has no escaping, so a value that contains -# `` leaves its call unparsed (`to_tool_calls` logs the dropped call). +# The body may hold only parameter blocks, so a value can contain any text except +# `` (unescapable) and a broken call can't run into the next one. +# `` is optional, for truncated output. _XML_TOOL_CALL_RE = re.compile( r"\s*\n]+)>\s*" r"((?:\n]+>(?:(?!).)*\s*)*)" @@ -489,14 +486,9 @@ def find_func(d: object) -> tuple[str | None, Mapping | None]: def _json_tool_call_spans(text: str) -> list[tuple[int, int]]: - """Return the `(start, end)` offsets of the JSON tool calls in `text`. - - Accepts what `json_extraction` and `find_func` accept, but on the raw text, - so the offsets line up with it. - """ + """Return the `(start, end)` offsets in `text` of the calls `json_extraction` would find.""" spans: list[tuple[int, int]] = [] - # `json_extraction` runs after whitespace is collapsed; this runs on the raw - # text, so it has to accept literal newlines inside strings. + # Raw text, unlike `json_extraction`'s, so strings may hold literal newlines. decoder = json.JSONDecoder(strict=False) index = text.find("{") while index != -1: @@ -513,14 +505,12 @@ def _json_tool_call_spans(text: str) -> list[tuple[int, int]]: def _parse_xml_tool_calls(llm_response: str) -> tuple[list[tuple[str, Mapping]], str]: - """Extract tool calls written in the XML function format. + """Extract XML-format tool calls, skipping any quoted in a JSON call's arguments. - Values are returned as raw strings, minus the single newline the template - wraps each one in; `to_tool_calls` coerces them against the tool schema. - Markup inside the arguments of a JSON tool call is data and is skipped. + Values are raw strings, minus the newline the template wraps each one in. Returns: - The calls, and `llm_response` with those calls blanked out. + The calls, and `llm_response` with them blanked out. """ matches = list(_XML_TOOL_CALL_RE.finditer(llm_response)) if not matches: @@ -549,12 +539,11 @@ def _parse_xml_tool_calls(llm_response: str) -> tuple[list[tuple[str, Mapping]], 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. - Two formats are recognized: JSON objects of the form - `{"name": ..., "arguments": {...}}` (bare, or inside `` tags), and - the XML function format that Granite 4.2 chat templates prescribe, - `VALUE` inside - `` tags. XML parameter values are returned as strings. XML - markup that appears inside a JSON call's arguments is treated as data. + Recognizes JSON calls (`{"name": ..., "arguments": {...}}`, bare or in + `` tags) and Granite 4.2's XML format + (`VALUE`), + 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. @@ -568,14 +557,13 @@ def parse_tools(llm_response: str) -> list[tuple[str, Mapping]]: def _parse_tool_calls(llm_response: str) -> tuple[list[tuple[str, Mapping, bool]], int]: - """Parse tool calls as `parse_tools` does, also saying which came from XML. + """Like `parse_tools`, but also flags XML calls and counts unparsed `` tags. Returns: - The calls as `(tool_name, arguments, from_xml)`, XML calls first, and - the number of `` tags that open none of them. + `(tool_name, arguments, from_xml)` per call, XML first, and the number + of `` tags that produced no call. """ - # The XML calls come back blanked out of `remainder`, so a JSON-looking - # parameter value is not parsed a second time as a call of its own. + # 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()) @@ -588,11 +576,9 @@ def _parse_tool_calls(llm_response: str) -> tuple[list[tuple[str, Mapping, bool] def _count_unparsed_tool_call_tags(text: str) -> int: - """Count the `` tags in `text` that don't open a JSON tool call. + """Count `` tags that open no JSON call, in text with XML calls blanked out. - Run on the output with the parsed XML calls blanked out. A tag quoted inside - a JSON call's arguments is data, and one directly followed by a JSON call - (or a JSON list of calls) is that call's wrapper; neither is counted. + Tags quoted in a JSON call's arguments, or wrapping a JSON call or list, don't count. """ if "" not in text: return 0 diff --git a/mellea/backends/utils.py b/mellea/backends/utils.py index f62c944d40..555ae51e38 100644 --- a/mellea/backends/utils.py +++ b/mellea/backends/utils.py @@ -148,13 +148,11 @@ def to_chat( def _decode_text_args( args: Mapping[str, Any], properties: Mapping[str, Any], required: Sequence[str] ) -> dict[str, Any]: - """Decode argument values that arrive as text but stand for something else. + """Decode XML-call values that stand for a list, dict or None. - Formats that carry every value as text (Granite's XML tool calls) render - list and dict arguments with `tojson` and a Python `None` as `None`. So a - string value for an object/array parameter is decoded as JSON, and `None` - or `null` becomes `None` for an optional parameter that can take it. Other - scalars are left for `validate_tool_arguments` to coerce. + 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. @@ -169,15 +167,13 @@ def _decode_text_args( if not isinstance(value, str): continue schema = properties.get(name) or {} - # `type` may be a list (`["array", "null"]`), at the top level or in an `anyOf` branch. + # `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"): - # Only where None is a valid value: the schema allows null, or the - # parameter is optional with no default other than None (the emitted - # schema leaves a None default out). + # 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 @@ -186,7 +182,7 @@ def _decode_text_args( try: parsed = json.loads(value) except json.JSONDecodeError: - # Leave it as-is; validate_tool_arguments reports the mismatch. + # Left for validate_tool_arguments to report. continue if isinstance(parsed, dict | list): decoded[name] = parsed @@ -227,7 +223,7 @@ def to_tool_calls( param_map = parameters["properties"] if len(param_map) == 0: tool_args = {} - elif from_xml: # XML calls carry every value as text; JSON calls are typed + elif from_xml: # only XML values arrive as text tool_args = _decode_text_args( tool_args, param_map, parameters.get("required") or [] ) diff --git a/test/backends/_granite_tokenizer.py b/test/backends/_granite_tokenizer.py index 1104ddfb15..d015350c11 100644 --- a/test/backends/_granite_tokenizer.py +++ b/test/backends/_granite_tokenizer.py @@ -1,12 +1,7 @@ # Copyright IBM Corp. All Rights Reserved. # SPDX-License-Identifier: Apache-2.0 -"""Load a cached Granite tokenizer, for tests that render the real chat template. - -Kept apart from the HF backend tests so that a light test file can use the -template without importing torch or `mellea.backends.huggingface`. Only the -tokenizer files are loaded; no model weights and no GPU. -""" +"""Load a cached Granite tokenizer, so tests can render the real chat template without torch or model weights.""" import pytest From a1f4a1ea75dae0cffa1863475a6d44bc3e9741a2 Mon Sep 17 00:00:00 2001 From: Nigel Jones Date: Fri, 2 Oct 2026 10:21:19 +0100 Subject: [PATCH 8/8] fix(backends): coerce the tool arguments that validate when one fails Lenient validate_tool_arguments returned every argument raw when any one failed, so a single bad value, including the explicit None an XML call decodes for an optional parameter, left ints and bools as strings. Keep the failing arguments as given and validate the rest. Addresses review on #1692. Assisted-by: Claude Code Signed-off-by: Nigel Jones --- mellea/backends/tools.py | 18 +++++++-- .../backends/test_pydantic_tool_parameters.py | 4 +- .../test_tool_validation_integration.py | 38 +++++++++++++++++-- test/backends/test_utils.py | 14 +++++-- 4 files changed, 61 insertions(+), 13 deletions(-) diff --git a/mellea/backends/tools.py b/mellea/backends/tools.py index e1884be529..362b7a89ff 100644 --- a/mellea/backends/tools.py +++ b/mellea/backends/tools.py @@ -613,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 @@ -896,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 diff --git a/test/backends/test_pydantic_tool_parameters.py b/test/backends/test_pydantic_tool_parameters.py index ec2a36d54b..1b3405b021 100644 --- a/test/backends/test_pydantic_tool_parameters.py +++ b/test/backends/test_pydantic_tool_parameters.py @@ -243,7 +243,7 @@ def test_missing_required_nested_field(self): # Missing 'body' field args = {"email": {"to": "user@example.com", "subject": "Test"}} - # In lenient mode, should return original args + # In lenient mode, the failing argument comes back as given validated = validate_tool_arguments(tool, args, strict=False) assert validated == args @@ -362,7 +362,7 @@ def test_flat_dict_instead_of_nested(self): with pytest.raises(ValidationError): validate_tool_arguments(tool, args, strict=True) - # In lenient mode, returns original + # In lenient mode, the failing and extra arguments come back as given validated = validate_tool_arguments(tool, args, strict=False) assert validated == args diff --git a/test/backends/test_tool_validation_integration.py b/test/backends/test_tool_validation_integration.py index a2dcadf337..be914457e6 100644 --- a/test/backends/test_tool_validation_integration.py +++ b/test/backends/test_tool_validation_integration.py @@ -156,15 +156,47 @@ class TestValidationModes: """Test strict vs. lenient validation modes.""" def test_lenient_mode_with_invalid_type(self): - """Test that lenient mode returns original args on validation failure.""" + """Test that lenient mode returns an invalid argument as given.""" args = {"name": "Test", "age": "not_a_number", "score": 95.5, "active": True} tool = MelleaTool.from_callable(typed_tool) validated = validate_tool_arguments(tool, args, strict=False) - # Should return original args + # The invalid argument comes back as given; the rest were already typed assert validated == args assert validated["age"] == "not_a_number" + def test_lenient_mode_coerces_the_arguments_that_validate(self): + """One invalid argument is passed through as given; the rest are still coerced.""" + args = { + "name": "Test", + "age": "not_a_number", + "score": "95.5", + "active": "true", + } + tool = MelleaTool.from_callable(typed_tool) + validated = validate_tool_arguments(tool, args, strict=False) + + assert validated == { + "name": "Test", + "age": "not_a_number", + "score": 95.5, + "active": True, + } + + def test_lenient_mode_keeps_unknown_args_when_another_fails(self): + """Extra arguments survive the per-argument fallback, as on success.""" + args = {"name": "Test", "age": "x", "score": "1", "active": "1", "extra": "e"} + tool = MelleaTool.from_callable(typed_tool) + validated = validate_tool_arguments(tool, args, strict=False) + + assert validated == { + "name": "Test", + "age": "x", + "score": 1.0, + "active": True, + "extra": "e", + } + def test_strict_mode_with_invalid_type(self): """Test that strict mode raises ValidationError on failure.""" args = {"name": "Test", "age": "not_a_number", "score": 95.5, "active": True} @@ -179,7 +211,7 @@ def test_lenient_mode_with_missing_required(self): tool = MelleaTool.from_callable(optional_tool) validated = validate_tool_arguments(tool, args, strict=False) - # Should return original args + # The missing argument stays missing; the rest come back validated assert validated == args def test_strict_mode_with_missing_required(self): diff --git a/test/backends/test_utils.py b/test/backends/test_utils.py index fd1fcc0283..44fbe3d90e 100644 --- a/test/backends/test_utils.py +++ b/test/backends/test_utils.py @@ -219,12 +219,16 @@ def test_decode_text_args_list_type_inside_any_of(): def test_to_tool_calls_granite_xml_none_text_for_optional_param_is_null(text: str): """The template renders a Python `None` as `None`; for an optional parameter that is a null.""" - def search(query: str, limit: int | None = None) -> str: + def search( + query: str, page: int, limit: int | None = None, exact: bool = False + ) -> str: """Search. Args: query: what to search for + page: the results page limit: optional result limit + exact: match the query exactly """ return query @@ -233,11 +237,13 @@ def search(query: str, limit: int | None = None) -> str: registry: dict[str, AbstractMelleaTool] = { "search": MelleaTool.from_callable(search) } - raw = _tool_call_xml("search", {"query": "cats", "limit": text}) + raw = _tool_call_xml( + "search", {"query": "cats", "page": "2", "limit": text, "exact": "True"} + ) result = to_tool_calls(registry, raw) assert result is not None - assert result[0].args["query"] == "cats" - assert result[0].args["limit"] is None + # The null must not stop the other arguments being coerced. + assert result[0].args == {"query": "cats", "page": 2, "limit": None, "exact": True} def test_to_tool_calls_none_text_for_required_param_is_kept():