From 3ee2baf71dea039c073a038e645f2e2919852ccc Mon Sep 17 00:00:00 2001
From: George Hotz <72895+geohot@users.noreply.github.com>
Date: Fri, 17 Jul 2026 10:20:01 -0700
Subject: [PATCH] llm: add tool calling support (kimi) (#17061)
* llm: add tool calling support
* simpler
* cls
* gpt cleanup
* more gpt cleanups
* tests for tools calling
---
test/null/test_llm_server.py | 92 +++++++++++++++++++++++++++++++++++-
tinygrad/llm/cli.py | 89 +++++++++++++++++++++++++++++++---
2 files changed, 173 insertions(+), 8 deletions(-)
diff --git a/test/null/test_llm_server.py b/test/null/test_llm_server.py
index 410b9daa84..87a7e28d38 100644
--- a/test/null/test_llm_server.py
+++ b/test/null/test_llm_server.py
@@ -1,4 +1,4 @@
-import unittest, threading, time
+import unittest, threading, time, json
from unittest.mock import Mock
class TestLLMServer(unittest.TestCase):
@@ -156,5 +156,95 @@ class TestLLMServer(unittest.TestCase):
self.assertEqual(data["data"][0]["id"], "test-model")
self.assertEqual(data["data"][0]["object"], "model")
+class TestLLMToolCalls(unittest.TestCase):
+ """Tool calling through the OpenAI-compatible HTTP API."""
+
+ @classmethod
+ def setUpClass(cls):
+ cls.mock_tok = Mock()
+ cls.mock_tok.encode = Mock(return_value=[200, 201, 202])
+ cls.mock_tok.decode = Mock(return_value="")
+ cls.mock_tok.preset = "qwen2"
+ cls.mock_tok.bos_id, cls.mock_tok.eos_id, cls.mock_tok.eot_id = None, 999, None
+ cls.mock_tok.is_end = Mock(return_value=False)
+
+ cls.mock_model = Mock()
+ cls.mock_model.get_start_pos = Mock(return_value=0)
+
+ from tinygrad.llm.cli import LLMServer
+ import jinja2
+ # .items() matches tool-aware templates and ensures OpenAI JSON argument strings are normalized before rendering the next turn.
+ template = jinja2.Template("""{% for m in messages %}{{ m.content or '' }}{% for tc in m.tool_calls or [] %}
+ {% for key, value in tc.function.arguments.items() %}{{ key }}={{ value }}{% endfor %}{% endfor %}{% endfor %}""")
+ cls.server = LLMServer(('127.0.0.1', 0), cls.mock_model, "tool-model", cls.mock_tok, template)
+ cls.port = cls.server.server_address[1]
+ cls.server_thread = threading.Thread(target=cls.server.serve_forever, daemon=True)
+ cls.server_thread.start()
+ time.sleep(0.1)
+
+ from openai import OpenAI
+ cls.client = OpenAI(base_url=f"http://127.0.0.1:{cls.port}/v1", api_key="test")
+
+ @classmethod
+ def tearDownClass(cls):
+ cls.server.shutdown()
+ cls.server.server_close()
+
+ def set_output(self, text:str):
+ pieces = dict(enumerate(text, 1))
+ self.mock_tok.stream_decoder = Mock(return_value=lambda tid=None: pieces[tid] if tid is not None else "")
+ self.mock_model.generate = Mock(side_effect=lambda ids, **kwargs: iter(pieces))
+
+ @staticmethod
+ def tools():
+ return [{"type":"function", "function":{"name":"read", "description":"Read a file",
+ "parameters":{"type":"object", "properties":{"path":{"type":"string"}}, "required":["path"]}}}]
+
+ def test_streaming_tool_call(self):
+ self.set_output('before{"name":"read","arguments":{"path":"README.md"}}')
+ chunks = list(self.client.chat.completions.create(model="tool-model", messages=[{"role":"user", "content":"Read README.md"}],
+ tools=self.tools(), stream=True))
+ self.assertEqual("".join(c.choices[0].delta.content or "" for c in chunks if c.choices), "before")
+ calls = [tc for c in chunks if c.choices for tc in c.choices[0].delta.tool_calls or []]
+ self.assertEqual(len(calls), 1)
+ self.assertEqual(calls[0].function.name, "read")
+ self.assertEqual(json.loads(calls[0].function.arguments), {"path":"README.md"})
+ self.assertEqual(chunks[-1].choices[0].finish_reason, "tool_calls")
+
+ def test_multiple_xml_tool_calls(self):
+ self.set_output("\"a\""
+ "\"b\"")
+ response = self.client.chat.completions.create(model="tool-model", messages=[{"role":"user", "content":"Read a and b"}],
+ tools=self.tools())
+ self.assertEqual([json.loads(tc.function.arguments)["path"] for tc in response.choices[0].message.tool_calls], ["a", "b"])
+ self.assertEqual(response.choices[0].finish_reason, "tool_calls")
+
+ def test_invalid_tool_call_becomes_content(self):
+ self.set_output("not a call")
+ response = self.client.chat.completions.create(model="tool-model", messages=[{"role":"user", "content":"Hello"}], tools=self.tools())
+ self.assertEqual(response.choices[0].message.content, "not a call")
+ self.assertIsNone(response.choices[0].message.tool_calls)
+ self.assertEqual(response.choices[0].finish_reason, "stop")
+
+ def test_tool_call_in_reasoning_is_not_executed(self):
+ self.set_output('draft {"name":"wrong","arguments":{}}answer')
+ response = self.client.chat.completions.create(model="tool-model", messages=[{"role":"user", "content":"Hello"}], tools=self.tools())
+ self.assertEqual(response.choices[0].message.content, "answer")
+ self.assertIsNone(response.choices[0].message.tool_calls)
+ self.assertEqual(response.choices[0].finish_reason, "stop")
+
+ def test_tool_result_round_trip(self):
+ self.set_output('{"name":"read","arguments":{"path":"README.md"}}')
+ first = self.client.chat.completions.create(model="tool-model", messages=[{"role":"user", "content":"Read README.md"}], tools=self.tools())
+ call = first.choices[0].message.tool_calls[0]
+ self.set_output("done")
+ second = self.client.chat.completions.create(model="tool-model", messages=[
+ {"role":"user", "content":"Read README.md"},
+ {"role":"assistant", "content":None, "tool_calls":[call.model_dump()]},
+ {"role":"tool", "tool_call_id":call.id, "content":"file contents"},
+ ], tools=self.tools())
+ self.assertEqual(second.choices[0].message.content, "done")
+ self.assertEqual(second.choices[0].finish_reason, "stop")
+
if __name__ == '__main__':
unittest.main()
diff --git a/tinygrad/llm/cli.py b/tinygrad/llm/cli.py
index c8d597d583..df67c7ecb0 100644
--- a/tinygrad/llm/cli.py
+++ b/tinygrad/llm/cli.py
@@ -9,6 +9,58 @@ from tinygrad.llm.model import Transformer
if TYPE_CHECKING:
import jinja2
+def parse_tool_call(s:str) -> tuple[str, typing.Any]|None:
+ s = s.strip()
+ if s.startswith("{"): # hermes JSON format: {"name": ..., "arguments": {...}}
+ try:
+ call = json.loads(s)
+ return call["name"], call.get("arguments", call.get("parameters", {}))
+ except (json.JSONDecodeError, KeyError): return None
+ # XML format: \n\nvalue\n...
+ if (fm := re.match(r"]+)>\s*(.*?)\s*(?:)?$", s, re.DOTALL)):
+ args = {}
+ for pm in re.finditer(r"]+)>\s*(.*?)\s*", fm.group(2), re.DOTALL):
+ try: args[pm.group(1)] = json.loads(pm.group(2))
+ except json.JSONDecodeError: args[pm.group(1)] = pm.group(2)
+ return fm.group(1), args
+ return None
+
+def normalize_messages(messages:list[dict]) -> None:
+ # chat templates expect tool_call arguments as dicts (OpenAI clients send JSON strings)
+ for m in messages:
+ for tc in m.get("tool_calls") or []:
+ if "function" in tc and isinstance(args := tc["function"].get("arguments"), str):
+ try: tc["function"]["arguments"] = json.loads(args)
+ except json.JSONDecodeError: pass
+
+class StreamRouter:
+ # routes streamed output text to (field, text) deltas, keeping tool_call regions in .buf for the final parse
+ def __init__(self):
+ self.buf = ""
+ self.mode = "undecided" # output inside a think block is sent as reasoning_content
+ def split(self, tag:str, final:bool) -> tuple[str, bool]:
+ # split buf on the first full tag, holding back a partial tag at the end unless final
+ if tag in self.buf:
+ before, self.buf = self.buf.split(tag, 1)
+ return before, True
+ hold = max((i for i in range(1, min(len(self.buf), len(tag))+1) if tag.startswith(self.buf[-i:])), default=0) if not final else 0
+ emit, self.buf = self.buf[:len(self.buf)-hold], self.buf[len(self.buf)-hold:]
+ return emit, False
+ def route(self, piece:str, final:bool=False) -> typing.Iterator[tuple[str, str]]:
+ self.buf += piece
+ if self.mode == "undecided": # decide whether the output starts with a think block
+ if not final and len(self.buf) < len("") and "".startswith(self.buf): return
+ self.mode, self.buf = ("reasoning", self.buf[len(""):]) if self.buf.startswith("") else ("content", self.buf)
+ if self.mode == "reasoning":
+ emit, done = self.split("", final)
+ if emit: yield "reasoning_content", emit
+ if not done: return
+ self.mode = "content"
+ if self.mode == "tool": return
+ emit, found = self.split("", final)
+ if emit: yield "content", emit
+ if found: self.mode, self.buf = "tool", "" + self.buf
+
class SimpleTokenizer:
def __init__(self, normal_tokens:dict[str, int], special_tokens:dict[str, int], preset:str="llama3",
bos_id:int|None=None, eos_id:int=0, eot_id:int|None=None):
@@ -137,20 +189,34 @@ class Handler(HTTPRequestHandler):
stderr_log(f"{self.path} {colored('--', 'BLACK')} "
f"in:{colored(f'{cache_start_pos:5d}', 'green')} +{len(ids)-cache_start_pos:5d} {colored('--', 'BLACK')} ")
tmpl = {"id":f"chatcmpl-{uuid.uuid4().hex[:24]}", "object":"chat.completion.chunk", "created":int(time.time()), "model":model_name}
- yield {"choices": [{"index":0, "delta":{"role":"assistant","content":""}, "finish_reason":None}], **tmpl}
+ def chunk(d:dict): return {"choices": [{"index":0, "delta":d, "finish_reason":None}], **tmpl}
+ yield chunk({"role":"assistant", "content":""})
out: list[int] = []
finish_reason = "stop"
st = time.perf_counter()
dec = tok.stream_decoder()
+ router = StreamRouter()
for next_id in model.generate(ids, temperature=temperature):
if len(out) == 0: stderr_log(f"prefill:{(len(ids)-cache_start_pos)/((pt:=time.perf_counter())-st):4.0f} tok/s {colored('--', 'BLACK')} ")
if tok.is_end(next_id): break
out.append(next_id)
- yield {"choices": [{"index":0, "delta":{"content":dec(next_id)}, "finish_reason":None}], **tmpl}
+ for field, delta in router.route(dec(next_id)): yield chunk({field:delta})
if max_tokens is not None and len(out) >= max_tokens:
finish_reason = "length"
break
- if (tail := dec()): yield {"choices": [{"index":0, "delta":{"content":tail}, "finish_reason":None}], **tmpl}
+ for field, delta in router.route(dec(), final=True): yield chunk({field:delta})
+ tool_calls: list[dict] = []
+ for m in re.finditer(r"\s*(.*?)\s*(?:|$)", router.buf, re.DOTALL):
+ if (parsed := parse_tool_call(m.group(1))) is None:
+ stderr_log(f"failed to parse tool call: {m.group(1)[:200]}")
+ yield chunk({"content":m.group(0)}) # don't silently drop output the client can't use
+ else:
+ name, args = parsed
+ tool_calls.append({"index":len(tool_calls), "id":f"call_{uuid.uuid4().hex[:24]}", "type":"function",
+ "function":{"name":name, "arguments":args if isinstance(args, str) else json.dumps(args)}})
+ if tool_calls:
+ yield chunk({"tool_calls":tool_calls})
+ if finish_reason == "stop": finish_reason = "tool_calls"
yield {"choices": [{"index":0, "delta":{},"finish_reason":finish_reason}], **tmpl}
if include_usage:
yield {"choices": [], "usage": {"prompt_tokens": len(ids), "completion_tokens": len(out), "total_tokens": len(ids) + len(out)}, **tmpl}
@@ -164,6 +230,7 @@ class Handler(HTTPRequestHandler):
if DEBUG >= 1: print(json.dumps(body, indent=2))
if self.path == "/v1/chat/completions":
# render and tokenize
+ if not isinstance(self.server.template, FallbackTemplate): normalize_messages(body["messages"])
rendered = self.server.template.render(messages=body["messages"], tools=body.get("tools"), add_generation_prompt=True)
ids: list[int] = self.server.tok.encode(rendered)
@@ -173,12 +240,20 @@ class Handler(HTTPRequestHandler):
max_tokens=max_tokens, temperature=float(body.get("temperature", 0.0)))
if body.get("stream"): self.stream_json(chunks)
else:
- out, finish_reason = [], "stop"
+ out, reasoning, tool_calls, finish_reason = [], [], [], "stop"
for c in chunks:
- if c["choices"] and c["choices"][0].get("delta", {}).get("content"): out.append(c["choices"][0]["delta"]["content"])
- if c["choices"] and c["choices"][0].get("finish_reason"): finish_reason = c["choices"][0]["finish_reason"]
+ if not c["choices"]: continue
+ choice = c["choices"][0]
+ if (delta := choice.get("delta", {})):
+ if delta.get("content"): out.append(delta["content"])
+ if delta.get("reasoning_content"): reasoning.append(delta["reasoning_content"])
+ tool_calls += [{k:v for k, v in tc.items() if k != "index"} for tc in delta.get("tool_calls", [])]
+ if choice.get("finish_reason"): finish_reason = choice["finish_reason"]
+ message: dict[str, typing.Any] = {"role":"assistant", "content":"".join(out) or None}
+ if reasoning: message["reasoning_content"] = "".join(reasoning)
+ if tool_calls: message["tool_calls"] = tool_calls
self.send_data(json.dumps({**c, "object":"chat.completion",
- "choices":[{"index":0, "message":{"role":"assistant","content":"".join(out)}, "finish_reason":finish_reason}]}).encode())
+ "choices":[{"index":0, "message":message, "finish_reason":finish_reason}]}).encode())
else:
raise RuntimeError(f"unhandled path {self.path}")