diff --git a/mellea/backends/tools.py b/mellea/backends/tools.py
index ae461dfb6..362b7a89f 100644
--- a/mellea/backends/tools.py
+++ b/mellea/backends/tools.py
@@ -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:
+# \n\n\nVALUE\n\n\n
+# 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*)*)"
+ r"\s*(?:)?",
+ re.DOTALL,
+)
+_XML_PARAMETER_RE = re.compile(r"\n]+)>(.*?)", 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
+ `` 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.
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 `` tags.
+
+ Returns:
+ `(tool_name, arguments, from_xml)` per call, XML first, and the number
+ of `` 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 `` 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 "" 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(
@@ -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
@@ -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
diff --git a/mellea/backends/utils.py b/mellea/backends/utils.py
index d9fbaa0b9..555ae51e3 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, Sequence
from typing import Any
from ..core import Context, MelleaLogger, ModelToolCall, Span
@@ -21,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
@@ -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
+ 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:
@@ -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} 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(
@@ -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 []
+ )
# Validate and coerce argument types
validated_args = validate_tool_arguments(func, tool_args, strict=False)
diff --git a/test/backends/_granite_tokenizer.py b/test/backends/_granite_tokenizer.py
new file mode 100644
index 000000000..d015350c1
--- /dev/null
+++ b/test/backends/_granite_tokenizer.py
@@ -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
diff --git a/test/backends/test_huggingface_filter_options.py b/test/backends/test_huggingface_filter_options.py
index 6aecdf010..f6af02561 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 d602bc2b5..47a7861a4 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,
)
@@ -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_pydantic_tool_parameters.py b/test/backends/test_pydantic_tool_parameters.py
index ec2a36d54..1b3405b02 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_helpers.py b/test/backends/test_tool_helpers.py
index a6c4bc611..0cf895ac2 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,138 @@ 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}}'})
+ ]
+
+
+_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_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"
+ 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'
+ assert parse_tools(raw) == [("get_weather", {"city": "Boston"})]
+
+
if __name__ == "__main__":
pytest.main([__file__])
diff --git a/test/backends/test_tool_validation_integration.py b/test/backends/test_tool_validation_integration.py
index a2dcadf33..be914457e 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 4889d1b68..44fbe3d90 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
@@ -15,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 ---
@@ -128,6 +133,300 @@ 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_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_text_args
+
+ properties = {"x": {"anyOf": [{"type": ["array", "null"]}]}}
+ 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, 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
+
+ from mellea.core.base import AbstractMelleaTool
+
+ registry: dict[str, AbstractMelleaTool] = {
+ "search": MelleaTool.from_callable(search)
+ }
+ raw = _tool_call_xml(
+ "search", {"query": "cats", "page": "2", "limit": text, "exact": "True"}
+ )
+ result = to_tool_calls(registry, raw)
+ assert result is not 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():
+ """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_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."""
+
+ 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)
+
+
+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,
+):
+ """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.
+
+ Loads the template only (no GPU, no model weights); skips if not locally cached.
+ """
+ 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 ---