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 ---