Add LongCat Flash tool parser (#810)
* Add LongCat Flash tool parser * Add unit tests for both xml and json formats
This commit is contained in:
@@ -291,7 +291,10 @@ class TokenizerWrapper:
|
||||
self._tool_call_end = tool_call_end
|
||||
|
||||
vocab = tokenizer.get_vocab()
|
||||
THINK_TOKENS = [("<think>", "</think>")]
|
||||
THINK_TOKENS = [
|
||||
("<think>", "</think>"),
|
||||
("<longcat_think>", "</longcat_think>"),
|
||||
]
|
||||
for think_start, think_end in THINK_TOKENS:
|
||||
if think_start in vocab and think_end in vocab:
|
||||
self._think_start = think_start
|
||||
@@ -472,6 +475,8 @@ def _infer_tool_parser(chat_template):
|
||||
return "minimax_m2"
|
||||
elif "<start_function_call>" in chat_template:
|
||||
return "function_gemma"
|
||||
elif "<longcat_tool_call>" in chat_template:
|
||||
return "longcat"
|
||||
elif "<arg_key>" in chat_template:
|
||||
return "glm47"
|
||||
elif "<tool_call>\n<function=" in chat_template:
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
# Copyright © 2026 Apple Inc.
|
||||
|
||||
import ast
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
import regex as re
|
||||
|
||||
_func_name_regex = re.compile(r"^(.*?)<longcat_arg_key>", re.DOTALL)
|
||||
_func_arg_regex = re.compile(
|
||||
r"<longcat_arg_key>(.*?)</longcat_arg_key>(?:\\n|\s)*<longcat_arg_value>(.*?)</longcat_arg_value>",
|
||||
re.DOTALL,
|
||||
)
|
||||
|
||||
tool_call_start = "<longcat_tool_call>"
|
||||
tool_call_end = "</longcat_tool_call>"
|
||||
|
||||
|
||||
def _is_string_type(
|
||||
tool_name: str,
|
||||
arg_name: str,
|
||||
tools: list[Any] | None,
|
||||
) -> bool:
|
||||
if tools is None:
|
||||
return False
|
||||
for tool in tools:
|
||||
func = tool["function"]
|
||||
if func["name"] == tool_name:
|
||||
params = func["parameters"]
|
||||
if params is None:
|
||||
return False
|
||||
arg_type = params.get("properties", {}).get(arg_name, {}).get("type", None)
|
||||
return arg_type == "string"
|
||||
return False
|
||||
|
||||
|
||||
def _deserialize(value: str) -> Any:
|
||||
try:
|
||||
return json.loads(value)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
try:
|
||||
return ast.literal_eval(value)
|
||||
except Exception:
|
||||
pass
|
||||
return value
|
||||
|
||||
|
||||
def parse_tool_call(text: str, tools: list[Any] | None = None):
|
||||
text = text.strip()
|
||||
|
||||
if text.startswith("{"):
|
||||
try:
|
||||
return json.loads(text)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
func_name = _func_name_regex.search(text).group(1).strip()
|
||||
pairs = _func_arg_regex.findall(text)
|
||||
arg_dct = {}
|
||||
for key, value in pairs:
|
||||
arg_key = key.strip()
|
||||
arg_val = value.strip()
|
||||
if not _is_string_type(func_name, arg_key, tools):
|
||||
arg_val = _deserialize(arg_val)
|
||||
arg_dct[arg_key] = arg_val
|
||||
return dict(name=func_name, arguments=arg_dct)
|
||||
@@ -6,6 +6,7 @@ from mlx_lm.tool_parsers import (
|
||||
glm47,
|
||||
json_tools,
|
||||
kimi_k2,
|
||||
longcat,
|
||||
minimax_m2,
|
||||
qwen3_coder,
|
||||
)
|
||||
@@ -14,12 +15,23 @@ from mlx_lm.tool_parsers import (
|
||||
class TestToolParsing(unittest.TestCase):
|
||||
|
||||
def test_parsers(self):
|
||||
parsers = [function_gemma, glm47, json_tools, kimi_k2, minimax_m2, qwen3_coder]
|
||||
parsers = [
|
||||
function_gemma,
|
||||
glm47,
|
||||
json_tools,
|
||||
longcat,
|
||||
longcat,
|
||||
kimi_k2,
|
||||
minimax_m2,
|
||||
qwen3_coder,
|
||||
]
|
||||
|
||||
test_cases = [
|
||||
"call:multiply{a:12234585,b:48838483920}",
|
||||
"multiply<arg_key>a</arg_key><arg_value>12234585</arg_value><arg_key>b</arg_key><arg_value>48838483920</arg_value>",
|
||||
'{"name": "multiply", "arguments": {"a": 12234585, "b": 48838483920}}',
|
||||
"multiply<longcat_arg_key>a</longcat_arg_key>\n<longcat_arg_value>12234585</longcat_arg_value>\n<longcat_arg_key>b</longcat_arg_key>\n<longcat_arg_value>48838483920</longcat_arg_value>",
|
||||
'{"name": "multiply", "arguments": {"a": 12234585, "b": 48838483920}}',
|
||||
'<|tool_call_begin|>functions.multiply:0<|tool_call_argument_begin|>{"a": 12234585, "b": 48838483920}<|tool_call_end|>',
|
||||
'<invoke name="multiply">\n<parameter name="a">12234585</parameter>\n<parameter name="b">48838483920</parameter>\n</invoke>',
|
||||
"<function=multiply>\n<parameter=a>\n12234585\n</parameter>\n<parameter=b>\n48838483920\n</parameter>\n</function>",
|
||||
@@ -55,6 +67,8 @@ class TestToolParsing(unittest.TestCase):
|
||||
"call:get_current_temperature{location:<escape>London<escape>}",
|
||||
'get_current_temperature<arg_key>location</arg_key><arg_value>"London"</arg_value>',
|
||||
'{"name": "get_current_temperature", "arguments": {"location": "London"}}',
|
||||
"get_current_temperature<longcat_arg_key>location</longcat_arg_key>\n<longcat_arg_value>London</longcat_arg_value>",
|
||||
'{"name": "get_current_temperature", "arguments": {"location": "London"}}',
|
||||
'<|tool_call_begin|>functions.get_current_temperature:0<|tool_call_argument_begin|>{"location": "London"}<|tool_call_end|>',
|
||||
'<invoke name="get_current_temperature">\n<parameter name="location">London</parameter>\n</invoke>',
|
||||
"<function=get_current_temperature>\n<parameter=location>\nLondon\n</parameter>\n</function>",
|
||||
|
||||
Reference in New Issue
Block a user