diff --git a/python/src/edge0/engine/ling.py b/python/src/edge0/engine/ling.py index 936803c..810ed5c 100644 --- a/python/src/edge0/engine/ling.py +++ b/python/src/edge0/engine/ling.py @@ -343,13 +343,15 @@ def _chat_template(self): ).from_string(src) return self._chat_tpl - def encode_chat(self, messages, think=None) -> list: + def encode_chat(self, messages, think=None, tools=None) -> list: """Tokenize chat messages with the deployment chat template. ``think`` mirrors THINK_MODE: True renders "detailed thinking on" (the model answers with a reasoning preamble), False renders "detailed thinking off" (direct answer). Defaults to the - engine's ``think`` flag. + engine's ``think`` flag. ``tools`` is the OpenAI-style function + list from the request, forwarded to the template's own ``# Tools`` + system-prompt section and ```` format instructions. """ if think is None: think = self.think @@ -366,8 +368,15 @@ def encode_chat(self, messages, think=None) -> list: messages=msgs, add_generation_prompt=True, enable_thinking=bool(think), - tools=None, + tools=tools, ) # transformers encode would prepend/append special tokens by # default; the template text is already complete. return self._tok.encode(text, add_special_tokens=False) + + def parse_tool_calls(self, text: str): + """Split generated ``text`` into (content, OpenAI tool_calls) per + this checkpoint's chat template ``name...`` + dialect (see edge0.server.tool_calls).""" + from edge0.server.tool_calls import parse_ling_tool_calls + return parse_ling_tool_calls(text) diff --git a/python/src/edge0/engine/qwen.py b/python/src/edge0/engine/qwen.py index b22a107..6f8df55 100644 --- a/python/src/edge0/engine/qwen.py +++ b/python/src/edge0/engine/qwen.py @@ -113,6 +113,13 @@ class Qwen35Engine(Edge0Engine): name = "edge0-35b" + def parse_tool_calls(self, text: str): + """Split generated ``text`` into (content, OpenAI tool_calls) per + this checkpoint's chat template ```` + dialect (see edge0.server.tool_calls).""" + from edge0.server.tool_calls import parse_qwen_tool_calls + return parse_qwen_tool_calls(text) + def __init__(self, model_dir: str, cfg, tokenizer=None): # NAN_BANG_COLLAPSE_FIX parity (engine/ling.py): hidden clip default # 1000 unless the deployer overrides QWEN_HIDDEN_CLIP explicitly. diff --git a/python/src/edge0/server/app.py b/python/src/edge0/server/app.py index 9d94b1b..fcbda96 100644 --- a/python/src/edge0/server/app.py +++ b/python/src/edge0/server/app.py @@ -28,9 +28,7 @@ def _error(status: int, msg: str) -> tuple: - if _HAS_FLASK: - return jsonify({"error": {"message": msg, "type": "invalid_request"}}), status - return (json.dumps({"error": {"message": msg}}), status) + return {"error": {"message": msg, "type": "invalid_request"}}, status def _split_think(text: str, think: bool): @@ -47,6 +45,21 @@ def _split_think(text: str, think: bool): return "", text +def _extract_tool_calls(server: QueueServer, req, content: str): + """(content, tool_calls) after stripping any ```` block the + engine's chat template asked the model to emit. ``tool_calls`` is None + when the engine has no parser (unsupported family) or the request + didn't ask for tools, or when no call was found in ``content``.""" + if req.tool_choice == "none" or (not req.tools + and req.tool_choice != "required"): + return content, None + parse = getattr(server.engine, "parse_tool_calls", None) + if not callable(parse): + return content, None + content, calls = parse(content) + return content, (calls or None) + + def _chat_once(server: QueueServer, payload: dict): req = parse_chat_request(payload) tokens, meta = server.chat(req) @@ -54,6 +67,14 @@ def _chat_once(server: QueueServer, payload: dict): think = bool(req.enable_thinking if req.enable_thinking is not None else getattr(server.engine, "think", False)) reasoning, content = _split_think(text, think) + content, tool_calls = _extract_tool_calls(server, req, content) + message = { + "role": "assistant", + "content": content, + "reasoning_content": reasoning, + } + if tool_calls: + message["tool_calls"] = tool_calls return { "id": f"chatcmpl-{int(time.time() * 1000)}", "object": "chat.completion", @@ -61,12 +82,8 @@ def _chat_once(server: QueueServer, payload: dict): "model": server.model_name, "choices": [{ "index": 0, - "message": { - "role": "assistant", - "content": content, - "reasoning_content": reasoning, - }, - "finish_reason": "stop", + "message": message, + "finish_reason": "tool_calls" if tool_calls else "stop", }], "usage": meta["usage"], } @@ -78,8 +95,20 @@ def _chat_stream(server: QueueServer, payload: dict): finished = object() request_id = f"chatcmpl-{int(time.time() * 1000)}" created = int(time.time()) + # Tool-call XML (...) must not leak into content + # deltas verbatim -- the whole point of #115 is that a client asking + # for tools gets a parsed message.tool_calls, not raw template markup. + # Detecting a block incrementally, token by token, needs a lookback + # buffer for a tag that can straddle token boundaries; simpler and + # just as correct: buffer the full response and parse once, same as + # the non-streaming path. Requests with no tools (the common case) + # keep the existing immediate per-token streaming, unchanged. + buffering = bool(req.tools) and req.tool_choice != "none" and callable( + getattr(server.engine, "parse_tool_calls", None)) def on_token(tid: int): + if buffering: + return text = decode_tokens(server.engine, [tid]) events.put(sse_format({ "id": request_id, "object": "chat.completion.chunk", @@ -91,12 +120,29 @@ def on_token(tid: int): def produce(): try: - _, meta = server.chat(req, on_token=on_token) + tokens, meta = server.chat(req, on_token=on_token) + finish_reason = "stop" + delta = {} + if buffering: + text = decode_tokens(server.engine, tokens) + think = bool( + req.enable_thinking if req.enable_thinking is not None + else getattr(server.engine, "think", False)) + _, content = _split_think(text, think) + content, tool_calls = _extract_tool_calls( + server, req, content) + if tool_calls: + delta = {"content": content, "tool_calls": [ + {"index": i, **call} for i, call in enumerate(tool_calls) + ]} + finish_reason = "tool_calls" + else: + delta = {"content": content} events.put(sse_format({ "id": request_id, "object": "chat.completion.chunk", "created": created, "model": server.model_name, - "choices": [{"index": 0, "delta": {}, - "finish_reason": "stop"}], + "choices": [{"index": 0, "delta": delta, + "finish_reason": finish_reason}], "usage": meta["usage"], }).encode("utf-8")) events.put(b"data: [DONE]\n\n") @@ -119,6 +165,11 @@ def build_app_handlers(server: QueueServer): """Return a handler dispatch dict shared by both transports.""" def handle_chat(payload: dict): + # Validate before sending streaming response headers. + try: + parse_chat_request(payload) + except ValueError as exc: + return _error(400, str(exc)) if payload.get("stream"): if _HAS_FLASK: return Response( @@ -171,8 +222,9 @@ def chat(): out = handlers["POST /v1/chat/completions"](payload) if isinstance(out, Response): return out - if isinstance(out, tuple) and out and isinstance(out[0], Response): - return out + if isinstance(out, tuple): + body, status = out + return jsonify(body), status return jsonify(out) @app.post("/v1/completions") @@ -180,7 +232,8 @@ def completions(): payload = request.get_json(force=True, silent=True) or {} out = handlers["POST /v1/completions"](payload) if isinstance(out, tuple): - return out + body, status = out + return jsonify(body), status return jsonify(out) return app @@ -237,18 +290,20 @@ def do_POST(self): # noqa: N802 return length = int(self.headers.get("Content-Length", 0)) payload = json.loads(self.rfile.read(length) or b"{}") + if path == "/v1/chat/completions": + try: + parse_chat_request(payload) + except ValueError as exc: + self._json(400, {"error": {"message": str(exc), + "type": "invalid_request"}}) + return if path == "/v1/chat/completions" and payload.get("stream"): self._sse(_chat_stream(self.server_q, payload)) return out = handlers[("POST " + path)](payload) if isinstance(out, tuple): - body, status, headers = out - self.send_response(status) - for k, v in headers.items(): - self.send_header(k, v) - self.send_header("Content-Length", str(len(body))) - self.end_headers() - self.wfile.write(body) + body, status = out + self._json(status, body) else: self._json(200, out) diff --git a/python/src/edge0/server/chat.py b/python/src/edge0/server/chat.py index a1502bd..1cdb93a 100644 --- a/python/src/edge0/server/chat.py +++ b/python/src/edge0/server/chat.py @@ -21,6 +21,9 @@ class ChatMessage: role: str content: str + tool_calls: list[dict] | None = None + tool_call_id: str | None = None + name: str | None = None @dataclass @@ -34,18 +37,50 @@ class ChatRequest: seed: int | None = None stream: bool = False enable_thinking: bool | None = None + tools: list[dict] | None = None + tool_choice: Any = None raw: dict[str, Any] = field(default_factory=dict) +def _template_tool_calls(calls) -> list[dict] | None: + if calls is None: + return None + if not isinstance(calls, list): + raise ValueError("tool_calls must be a list of function calls") + normalized = [] + for call in calls: + if not isinstance(call, dict) or not isinstance(call.get("function"), dict): + raise ValueError("tool_calls entries must contain a function object") + function = call["function"] + arguments = function.get("arguments", {}) + # OpenAI messages serialize arguments; checkpoint templates iterate them. + if isinstance(arguments, str): + try: + arguments = json.loads(arguments) + except ValueError as exc: + raise ValueError( + "tool_calls function arguments must be a JSON object") from exc + if not isinstance(arguments, dict): + raise ValueError("tool_calls function arguments must be a JSON object") + normalized.append({**call, "function": {**function, "arguments": arguments}}) + return normalized + + def parse_chat_request(payload: dict) -> ChatRequest: msgs = [] for m in payload.get("messages", []): role = str(m.get("role", "user")) - content = m.get("content", "") + content = m.get("content") or "" # tool-call-only assistant turns + # send content: null, not an omitted key if isinstance(content, list): # multi-part content: join text parts content = "".join( p.get("text", "") for p in content if isinstance(p, dict)) - msgs.append(ChatMessage(role=role, content=str(content))) + msgs.append(ChatMessage( + role=role, content=str(content), + tool_calls=_template_tool_calls(m.get("tool_calls")), + tool_call_id=m.get("tool_call_id"), + name=m.get("name"), + )) return ChatRequest( model=str(payload.get("model", "")), messages=msgs, @@ -56,6 +91,8 @@ def parse_chat_request(payload: dict) -> ChatRequest: seed=payload.get("seed"), stream=bool(payload.get("stream", False)), enable_thinking=payload.get("enable_thinking"), + tools=payload.get("tools"), + tool_choice=payload.get("tool_choice"), raw=payload, ) @@ -68,6 +105,16 @@ def __init__(self, engine: Edge0Engine, req: ChatRequest): self.req = req self._tok = engine._tok + def _tools(self) -> list[dict] | None: + """Tool definitions to forward to the chat template, or None if + the request has none or explicitly disabled calling + (``tool_choice: "none"``). Forcing a specific function or + ``tool_choice: "required"`` would need constrained decoding, + which isn't implemented; both currently behave like "auto".""" + if not self.req.tools or self.req.tool_choice == "none": + return None + return self.req.tools + def prompt_ids(self) -> list[int]: tok = self._tok # Families with a vendored chat template @@ -82,7 +129,8 @@ def prompt_ids(self) -> list[int]: if think is None: think = getattr(self.engine, "think", False) return encode( - [m.__dict__ for m in self.req.messages], think=bool(think)) + [m.__dict__ for m in self.req.messages], think=bool(think), + tools=self._tools()) if hasattr(tok, "apply_chat_template"): # The qwen35 template supports ``enable_thinking``: False # renders the canonical no-think prompt — an EMPTY think @@ -101,7 +149,7 @@ def prompt_ids(self) -> list[int]: try: text = tok.apply_chat_template( msgs, tokenize=False, add_generation_prompt=True, - enable_thinking=bool(think)) + enable_thinking=bool(think), tools=self._tools()) except TypeError: # tokenizer template without the kwarg: plain render try: diff --git a/python/src/edge0/server/tool_calls.py b/python/src/edge0/server/tool_calls.py new file mode 100644 index 0000000..0a0538f --- /dev/null +++ b/python/src/edge0/server/tool_calls.py @@ -0,0 +1,136 @@ +"""Parse the XML tool-call blocks each model family's chat template asks +it to emit, and convert them to the OpenAI ``message.tool_calls`` schema. + +The two supported checkpoints use two different XML dialects (confirmed +against the real ``chat_template.jinja`` shipped with each checkpoint, +not assumed): + +* **edge0-8b (Ling)** — ``{name}k + v...``, one ``arg_key``/``arg_value`` + pair per argument. +* **edge0-35b (Qwen3.5)** — ``\\n\\n + \\nvalue\\n\\n...\\n``. + +Both formats can appear more than once per turn (multiple tool calls in +one response); each ``...`` block becomes one +OpenAI tool-call entry with a synthesized ``id``. +""" + +from __future__ import annotations + +import json +import re +import uuid + +_TOOL_CALL_BLOCK = re.compile(r"(.*?)", re.DOTALL) + +_LING_ARG = re.compile( + r"(.*?)\s*(.*?)", re.DOTALL) + +_QWEN_FUNCTION = re.compile( + r"\s*]+)>\s*(.*?)\s*\s*", re.DOTALL) +_QWEN_PARAM = re.compile( + r"]+)>\n(.*?)\n", re.DOTALL) + + +def _coerce(raw: str): + """A bare arg value is a string unless it parses as JSON (numbers, + booleans, null, objects, arrays) -- mirrors the template's own + ``v if v is string else v|tojson`` split, in reverse.""" + try: + return json.loads(raw) + except (json.JSONDecodeError, ValueError): + return raw + + +def _tool_call_id() -> str: + return f"call_{uuid.uuid4().hex[:24]}" + + +def _to_openai(name: str, args: dict) -> dict: + return { + "id": _tool_call_id(), + "type": "function", + "function": {"name": name.strip(), "arguments": json.dumps(args)}, + } + + +def _split(text: str, parse_block) -> tuple[str | None, list[dict]]: + """Shared block extraction: find every ``...``, + remove it from the content, and parse its body with ``parse_block``. + Blocks ``parse_block`` cannot make sense of are left in ``content`` + rather than silently dropped.""" + calls = [] + leftover_spans = [] + pos = 0 + for m in _TOOL_CALL_BLOCK.finditer(text): + parsed = parse_block(m.group(1)) + if parsed is None: + continue + calls.append(_to_openai(*parsed)) + leftover_spans.append((pos, m.start())) + pos = m.end() + leftover_spans.append((pos, len(text))) + content = "".join(text[a:b] for a, b in leftover_spans).strip() + if calls and not content: + return None, calls + return content, calls + + +def _parse_arguments(text: str, pattern) -> dict | None: + args = {} + pos = 0 + for match in pattern.finditer(text): + if text[pos:match.start()].strip(): + return None + key = match.group(1).strip() + if not key: + return None + args[key] = _coerce(match.group(2)) + pos = match.end() + if text[pos:].strip(): + return None + return args + + +def _parse_ling_block(body: str): + name_match = re.match(r"^([^<]+)", body) + if not name_match: + return None + name = name_match.group(1).strip() + if not name: + return None + args = _parse_arguments(body[name_match.end():], _LING_ARG) + if args is None: + return None + return name, args + + +def parse_ling_tool_calls(text: str) -> tuple[str | None, list[dict]]: + """Split ``text`` into (remaining content, OpenAI tool_calls) for the + edge0-8b (Ling) XML dialect. Returns ``(text.strip(), [])`` when no + ```` block is present -- same stripping convention as + ``app._split_think``.""" + return _split(text, _parse_ling_block) + + +def _parse_qwen_block(body: str): + fn_match = _QWEN_FUNCTION.fullmatch(body) + if not fn_match: + return None + name = fn_match.group(1).strip() + if not name: + return None + params = fn_match.group(2) + args = _parse_arguments(params, _QWEN_PARAM) + if args is None: + return None + return name, args + + +def parse_qwen_tool_calls(text: str) -> tuple[str | None, list[dict]]: + """Split ``text`` into (remaining content, OpenAI tool_calls) for the + edge0-35b (Qwen3.5) XML dialect. Returns ``(text.strip(), [])`` when + no ```` block is present -- same stripping convention as + ``app._split_think``.""" + return _split(text, _parse_qwen_block) diff --git a/python/tests/fixtures/tool_templates/README.md b/python/tests/fixtures/tool_templates/README.md new file mode 100644 index 0000000..ef4346a --- /dev/null +++ b/python/tests/fixtures/tool_templates/README.md @@ -0,0 +1,8 @@ +These Apache-2.0 templates are copied from the official Edge0 checkpoints +for offline multi-turn tool-call regression tests. Model weights are not included. + +- Ling: https://huggingface.co/Edge0/Edge0-8B-A1B-preview/blob/269b9a2c4a69d897c50e9f4e125328481d7c0fcf/chat_template.jinja +- Qwen: https://huggingface.co/Edge0/Edge0-35B-A3B-preview/blob/3fe15cbf2bd5bbcdfd611035e1ac88971ff1cdab/chat_template.jinja + +The original model cards identify the license as Apache-2.0. Retain the +templates' original comments and consult the repository LICENSE. diff --git a/python/tests/fixtures/tool_templates/ling/chat_template.jinja b/python/tests/fixtures/tool_templates/ling/chat_template.jinja new file mode 100644 index 0000000..4be3ebd --- /dev/null +++ b/python/tests/fixtures/tool_templates/ling/chat_template.jinja @@ -0,0 +1,130 @@ +{#- Bailing V3 chat template -#} +{#- Supports: thinking option, tool calling -#} + +{#- ==================== thinking option normalization ==================== -#} +{%- if enable_thinking is defined %} +{%- if enable_thinking %} +{%- set thinking_option = 'on' %} +{%- else %} +{%- set thinking_option = 'off' %} +{%- endif %} +{%- elif thinking_option is not defined %} +{%- set thinking_option = 'on' %} +{%- endif %} + +{#- ==================== preserved thinking ==================== -#} +{% set preserved_thinking = true %} + +{#- ==================== system message ==================== -#} +{{- 'SYSTEM' }} +{%- if tools %} + {%- if messages[0].role == 'system' %} + {{- messages[0].content + '\n' }} + {%- endif %} + {{- "# Tools\n\nYou may call one or more functions to assist with the user query.\n\nYou are provided with function signatures within XML tags:\n" }} + {%- for tool in tools %} + {{- "\n" }} + {{- tool | tojson }} + {%- endfor %} + {{- "\n\n\nIf none of the functions can be used, point it out. If the given question lacks the parameters required by the function, also point it out.\nIf you need to use a function, for each function call, output the function name and arguments within the following XML format:\n{function-name}\n{arg-key-1}\n{arg-value-1}\n{arg-key-2}\n{arg-value-2}\n...\n\n" }} + {%- if messages[0].role == 'system' and messages[0].content is string and ('detailed thinking on' in messages[0].content or 'detailed thinking off' in messages[0].content) %} + {{- '<|role_end|>' }} + {%- else %} + {{- 'detailed thinking ' + thinking_option + '<|role_end|>' }} + {%- endif %} +{%- else %} + {%- if messages[0].role == 'system' %} + {%- if 'detailed thinking on' in messages[0].content or 'detailed thinking off' in messages[0].content %} + {{- messages[0].content + '<|role_end|>' }} + {%- else %} + {{- messages[0].content + '\n' }} + {{- 'detailed thinking ' + thinking_option + '<|role_end|>' }} + {%- endif %} + {% else %} + {{- 'detailed thinking ' + thinking_option + '<|role_end|>' }} + {%- endif %} +{%- endif %} +{%- set ns = namespace(multi_step_tool=true, last_query_index=messages|length - 1) %} +{%- for message in messages[::-1] %} + {%- set index = (messages|length - 1) - loop.index0 %} + {%- if ns.multi_step_tool and message.role == "user" and message.content is string and not(message.content.startswith('') and message.content.endswith('')) %} + {%- set ns.multi_step_tool = false %} + {%- set ns.last_query_index = index %} + {%- endif %} +{%- endfor %} +{%- for message in messages %} + {%- if message.content is string %} + {%- set content = message.content %} + {%- else %} + {%- set content = '' %} + {%- endif %} + {%- if message.role == "user" %} + {{- 'HUMAN' + message.content + '<|role_end|>' }} + {%- elif message.role == "system" and not loop.first %} + {{- 'SYSTEM' + message.content + '<|role_end|>' }} + {%- elif message.role == "assistant" %} + {%- set reasoning_content = '' %} + {%- if message.reasoning_content is string and message.reasoning_content != '' %} + {%- set reasoning_content = message.reasoning_content %} + {%- else %} + {%- if '' in content %} + {%- set reasoning_content = content.split('')[0].rstrip('\n').split('')[-1].lstrip('\n') %} + {%- set content = content.split('')[-1].lstrip('\n') %} + {%- endif %} + {%- endif %} + {%- if preserved_thinking or loop.index0 > ns.last_query_index %} + {%- if reasoning_content != '' %} + {{- 'ASSISTANT' + '\n' + reasoning_content.strip('\n') + '' + content.lstrip('\n') }} + {%- else %} + {{- 'ASSISTANT\n' + content }} + {%- endif %} + {%- else %} + {{- 'ASSISTANT\n' + content }} + {%- endif %} + {%- if message.tool_calls %} + {%- for tool_call in message.tool_calls %} + {%- if (loop.first and content) or (not loop.first) %} + {{- '\n' }} + {%- endif %} + {%- set tc = tool_call %} + {%- if tool_call.function %} + {%- set tc = tool_call.function %} + {%- endif %} + {{- '' + tc.name }} + {% set _args = tc.arguments %} + {%- for k, v in _args.items() %} + {{- '' + k + '' }} + {{- '\n' }} + {%- if v is string %} + {{- v }} + {%- else %} + {{- v | tojson(ensure_ascii=False) }} + {%- endif %} + {{- '' }} + {%- endfor %} + {{- '\n' }} + {%- endfor %} + {%- endif %} + {{- '<|role_end|>' }} + {%- elif message.role == "tool" %} + {%- if loop.first or (messages[loop.index0 - 1].role != "tool") %} + {{- 'OBSERVATION' }} + {%- endif %} + {{- '\n\n' }} + {{- content }} + {{- '\n' }} + {%- if loop.last or (messages[loop.index0 + 1].role != "tool") %} + {{- '<|role_end|>' }} + {%- endif %} + {%- endif %} +{%- endfor %} + +{#- ==================== generation prompt ==================== -#} +{%- if add_generation_prompt %} + {{- 'ASSISTANT' }} + {%- if thinking_option == 'on' %} + {{- '\n' }} + {%- elif thinking_option == 'off' %} + {{- '\n' }} + {%- endif %} +{%- endif %} diff --git a/python/tests/fixtures/tool_templates/qwen/chat_template.jinja b/python/tests/fixtures/tool_templates/qwen/chat_template.jinja new file mode 100644 index 0000000..4540e44 --- /dev/null +++ b/python/tests/fixtures/tool_templates/qwen/chat_template.jinja @@ -0,0 +1,154 @@ +{%- set image_count = namespace(value=0) %} +{%- set video_count = namespace(value=0) %} +{%- macro render_content(content, do_vision_count, is_system_content=false) %} + {%- if content is string %} + {{- content }} + {%- elif content is iterable and content is not mapping %} + {%- for item in content %} + {%- if 'image' in item or 'image_url' in item or item.type == 'image' %} + {%- if is_system_content %} + {{- raise_exception('System message cannot contain images.') }} + {%- endif %} + {%- if do_vision_count %} + {%- set image_count.value = image_count.value + 1 %} + {%- endif %} + {%- if add_vision_id %} + {{- 'Picture ' ~ image_count.value ~ ': ' }} + {%- endif %} + {{- '<|vision_start|><|image_pad|><|vision_end|>' }} + {%- elif 'video' in item or item.type == 'video' %} + {%- if is_system_content %} + {{- raise_exception('System message cannot contain videos.') }} + {%- endif %} + {%- if do_vision_count %} + {%- set video_count.value = video_count.value + 1 %} + {%- endif %} + {%- if add_vision_id %} + {{- 'Video ' ~ video_count.value ~ ': ' }} + {%- endif %} + {{- '<|vision_start|><|video_pad|><|vision_end|>' }} + {%- elif 'text' in item %} + {{- item.text }} + {%- else %} + {{- raise_exception('Unexpected item type in content.') }} + {%- endif %} + {%- endfor %} + {%- elif content is none or content is undefined %} + {{- '' }} + {%- else %} + {{- raise_exception('Unexpected content type.') }} + {%- endif %} +{%- endmacro %} +{%- if not messages %} + {{- raise_exception('No messages provided.') }} +{%- endif %} +{%- if tools and tools is iterable and tools is not mapping %} + {{- '<|im_start|>system\n' }} + {{- "# Tools\n\nYou have access to the following functions:\n\n" }} + {%- for tool in tools %} + {{- "\n" }} + {{- tool | tojson }} + {%- endfor %} + {{- "\n" }} + {{- '\n\nIf you choose to call a function ONLY reply in the following format with NO suffix:\n\n\n\n\nvalue_1\n\n\nThis is the value for the second parameter\nthat can span\nmultiple lines\n\n\n\n\n\nReminder:\n- Function calls MUST follow the specified format: an inner block must be nested within XML tags\n- Required parameters MUST be specified\n- You may provide optional reasoning for your function call in natural language BEFORE the function call, but NOT after\n- If there is no function call available, answer the question like normal with your current knowledge and do not tell the user about function calls\n' }} + {%- if messages[0].role == 'system' %} + {%- set content = render_content(messages[0].content, false, true)|trim %} + {%- if content %} + {{- '\n\n' + content }} + {%- endif %} + {%- endif %} + {{- '<|im_end|>\n' }} +{%- else %} + {%- if messages[0].role == 'system' %} + {%- set content = render_content(messages[0].content, false, true)|trim %} + {{- '<|im_start|>system\n' + content + '<|im_end|>\n' }} + {%- endif %} +{%- endif %} +{%- set ns = namespace(multi_step_tool=true, last_query_index=messages|length - 1) %} +{%- for message in messages[::-1] %} + {%- set index = (messages|length - 1) - loop.index0 %} + {%- if ns.multi_step_tool and message.role == "user" %} + {%- set content = render_content(message.content, false)|trim %} + {%- if not(content.startswith('') and content.endswith('')) %} + {%- set ns.multi_step_tool = false %} + {%- set ns.last_query_index = index %} + {%- endif %} + {%- endif %} +{%- endfor %} +{%- if ns.multi_step_tool %} + {{- raise_exception('No user query found in messages.') }} +{%- endif %} +{%- for message in messages %} + {%- set content = render_content(message.content, true)|trim %} + {%- if message.role == "system" %} + {%- if not loop.first %} + {{- raise_exception('System message must be at the beginning.') }} + {%- endif %} + {%- elif message.role == "user" %} + {{- '<|im_start|>' + message.role + '\n' + content + '<|im_end|>' + '\n' }} + {%- elif message.role == "assistant" %} + {%- set reasoning_content = '' %} + {%- if message.reasoning_content is string %} + {%- set reasoning_content = message.reasoning_content %} + {%- else %} + {%- if '' in content %} + {%- set reasoning_content = content.split('')[0].rstrip('\n').split('')[-1].lstrip('\n') %} + {%- set content = content.split('')[-1].lstrip('\n') %} + {%- endif %} + {%- endif %} + {%- set reasoning_content = reasoning_content|trim %} + {%- if (preserve_thinking is defined and preserve_thinking is true) or (loop.index0 > ns.last_query_index) %} + {{- '<|im_start|>' + message.role + '\n\n' + reasoning_content + '\n\n\n' + content }} + {%- else %} + {{- '<|im_start|>' + message.role + '\n' + content }} + {%- endif %} + {%- if message.tool_calls and message.tool_calls is iterable and message.tool_calls is not mapping %} + {%- for tool_call in message.tool_calls %} + {%- if tool_call.function is defined %} + {%- set tool_call = tool_call.function %} + {%- endif %} + {%- if loop.first %} + {%- if content|trim %} + {{- '\n\n\n\n' }} + {%- else %} + {{- '\n\n' }} + {%- endif %} + {%- else %} + {{- '\n\n\n' }} + {%- endif %} + {%- if tool_call.arguments is defined %} + {%- for args_name, args_value in tool_call.arguments|items %} + {{- '\n' }} + {%- set args_value = args_value | string if args_value is string else args_value | tojson | safe %} + {{- args_value }} + {{- '\n\n' }} + {%- endfor %} + {%- endif %} + {{- '\n' }} + {%- endfor %} + {%- endif %} + {{- '<|im_end|>\n' }} + {%- elif message.role == "tool" %} + {%- if loop.previtem and loop.previtem.role != "tool" %} + {{- '<|im_start|>user' }} + {%- endif %} + {{- '\n\n' }} + {{- content }} + {{- '\n' }} + {%- if not loop.last and loop.nextitem.role != "tool" %} + {{- '<|im_end|>\n' }} + {%- elif loop.last %} + {{- '<|im_end|>\n' }} + {%- endif %} + {%- else %} + {{- raise_exception('Unexpected message role.') }} + {%- endif %} +{%- endfor %} +{%- if add_generation_prompt %} + {{- '<|im_start|>assistant\n' }} + {%- if enable_thinking is defined and enable_thinking is false %} + {{- '\n\n\n\n' }} + {%- else %} + {{- '\n' }} + {%- endif %} +{%- endif %} diff --git a/python/tests/test_server.py b/python/tests/test_server.py index bad235b..c153da2 100644 --- a/python/tests/test_server.py +++ b/python/tests/test_server.py @@ -12,7 +12,9 @@ import threading import time from http.server import ThreadingHTTPServer -from types import SimpleNamespace +from pathlib import Path +from types import MethodType, SimpleNamespace +from urllib.error import HTTPError from urllib import request as urlrequest import pytest @@ -51,8 +53,10 @@ def __init__(self, ids=(7, 8, 9)): def apply_chat_template(self, messages, tokenize=False, add_generation_prompt=True, - enable_thinking=False): + enable_thinking=False, tools=None): text = "" + json.dumps(messages) + if tools: + text += "" + json.dumps(tools) if add_generation_prompt: text += "<|im_start|>assistant\n" return text @@ -113,6 +117,42 @@ def reset(self): self.pos = 0 +class ToolCallTok(FakeTok): + """Decodes generated tokens to a Ling-style ```` block, the + same shape the real edge0-8b template asks the model to emit.""" + + def decode(self, tokens): + return ("calculator\n" + "expression\n" + "2 + 2\n" + "") + + +class ToolCallEngine(FakeEngine): + """FakeEngine + the real Ling tool-call parser, so tests exercise the + actual edge0.server.tool_calls integration, not a re-implementation + of it.""" + + def __init__(self, tok=None): + super().__init__(tok=tok or ToolCallTok()) + + def parse_tool_calls(self, text: str): + from edge0.server.tool_calls import parse_ling_tool_calls + return parse_ling_tool_calls(text) + + +_TOOLS = [{ + "type": "function", + "function": { + "name": "calculator", + "description": "Calculate an arithmetic expression.", + "parameters": {"type": "object", + "properties": {"expression": {"type": "string"}}, + "required": ["expression"]}, + }, +}] + + def _req(**kw) -> ChatRequest: payload = { "model": "fake", @@ -155,6 +195,70 @@ def test_parse_stream_flag_and_sampling(): assert req.stream is True +def test_parse_preserves_tools_and_tool_choice(): + req = _req(tools=_TOOLS, tool_choice="auto") + assert req.tools == _TOOLS + assert req.tool_choice == "auto" + + +def test_parse_message_tool_calls_and_tool_role_roundtrip(): + req = parse_chat_request({"messages": [ + {"role": "user", "content": "compute 2+2"}, + {"role": "assistant", "content": None, "tool_calls": [ + {"id": "call_1", "type": "function", + "function": {"name": "calculator", + "arguments": '{"expression": "2 + 2"}'}}]}, + {"role": "tool", "tool_call_id": "call_1", "content": "4"}, + ]}) + assert req.messages[1].tool_calls[0]["function"]["name"] == "calculator" + assert req.messages[1].content == "" # null, not the string "None" + assert req.messages[2].tool_call_id == "call_1" + + +@pytest.mark.parametrize("arguments", ['{"expression": "2 + 2"}', + {"expression": "2 + 2"}]) +def test_parse_normalizes_tool_arguments_without_mutating_payload(arguments): + payload = {"messages": [{"role": "assistant", "content": None, + "tool_calls": [{ + "id": "call_1", "type": "function", + "function": {"name": "calculator", + "arguments": arguments}, + }]}]} + before = json.dumps(payload) + req = parse_chat_request(payload) + call = req.messages[0].tool_calls[0] + assert call["function"]["arguments"] == {"expression": "2 + 2"} + assert call["id"] == "call_1" + assert json.dumps(payload) == before + + +@pytest.mark.parametrize("arguments", ["{broken", "[]", "null", '"text"', + "42", [], None]) +def test_parse_rejects_invalid_tool_arguments(arguments): + with pytest.raises(ValueError, match="arguments.*JSON object"): + parse_chat_request({"messages": [{ + "role": "assistant", "content": None, + "tool_calls": [{"type": "function", "function": { + "name": "calculator", "arguments": arguments}}], + }]}) + + +@pytest.mark.parametrize("calls", ["bad", {}, [None], [{"function": "bad"}]]) +def test_parse_rejects_invalid_tool_call_shapes(calls): + with pytest.raises(ValueError, match="tool_calls"): + parse_chat_request({"messages": [{ + "role": "assistant", "content": None, "tool_calls": calls, + }]}) + + +def test_parse_tool_call_without_arguments_uses_empty_object(): + req = parse_chat_request({"messages": [{ + "role": "assistant", "content": None, + "tool_calls": [{"function": {"name": "get_time"}}], + }]}) + assert req.messages[0].tool_calls[0]["function"]["arguments"] == {} + + # ---- sse / decode --------------------------------------------------------- @@ -293,6 +397,202 @@ def test_chat_once_response_shape(): assert out["usage"]["completion_tokens"] == 3 +# ---- tool calls (#115) ----------------------------------------------------- + + +@pytest.fixture(params=["ling", "qwen"]) +def real_template_engine(request): + from tokenizers import Tokenizer + from tokenizers.models import WordLevel + from transformers import PreTrainedTokenizerFast + + family = request.param + template_dir = Path(__file__).parent / "fixtures" / "tool_templates" / family + template = (template_dir / "chat_template.jinja").read_text(encoding="utf-8") + tokenizer = PreTrainedTokenizerFast( + tokenizer_object=Tokenizer(WordLevel({"[UNK]": 0}, unk_token="[UNK]")), + unk_token="[UNK]", chat_template=template, + ) + + class TemplateTok(FakeTok): + output = ( + "calculatorexpression" + "2 + 2" + if family == "ling" else + "\n\n\n" + "2 + 2\n\n\n" + ) + + def apply_chat_template(self, *args, **kwargs): + return tokenizer.apply_chat_template(*args, **kwargs) + + def encode(self, text, **kwargs): + return super().encode(text) + + def decode(self, tokens): + return self.output + + engine = ToolCallEngine(tok=TemplateTok()) + if family == "ling": + from edge0.engine.ling import Ling8BEngine + engine.dir = str(template_dir) + engine.think = False + engine._chat_tpl = None + engine._chat_template = MethodType(Ling8BEngine._chat_template, engine) + engine.encode_chat = MethodType(Ling8BEngine.encode_chat, engine) + else: + from edge0.engine.qwen import Qwen35Engine + engine.parse_tool_calls = MethodType(Qwen35Engine.parse_tool_calls, engine) + return engine + + +@pytest.mark.parametrize("stream", [False, True]) +def test_tool_call_roundtrip_renders_real_template(real_template_engine, stream): + engine = real_template_engine + server = QueueServer(engine) + user = {"role": "user", "content": "Use the calculator for 2 + 2."} + payload = {"messages": [user], "tools": _TOOLS} + if stream: + frames = b"".join(_chat_stream(server, payload)).decode().split("\n\n") + chunk = json.loads(frames[-3][6:])["choices"][0] + assert chunk["finish_reason"] == "tool_calls" + assistant = {"role": "assistant", **chunk["delta"]} + # A client accumulates indexed SSE deltas into a message. + assistant["tool_calls"] = [ + {k: v for k, v in call.items() if k != "index"} + for call in assistant["tool_calls"] + ] + else: + choice = _chat_once(server, payload)["choices"][0] + assert choice["finish_reason"] == "tool_calls" + assistant = choice["message"] + call = assistant["tool_calls"][0] + assert isinstance(call["function"]["arguments"], str) + engine._tok.output = "The result is 4." + result = _chat_once(server, { + "messages": [user, assistant, { + "role": "tool", "tool_call_id": call["id"], + "name": "calculator", "content": "4", + }], "tools": _TOOLS, + }) + prompt = engine._tok.encoded_text + assert "\n4\n" in prompt + assert "" in prompt + assert "2 + 2" in prompt + assert result["choices"][0]["message"]["content"] == "The result is 4." + + +def test_chat_once_ignores_tool_call_text_without_tools_field(): + """No ``tools`` in the request -> the raw XML the fake + generation returns is left as plain content, exactly today's + (buggy) behavior -- this only changes when the client opts in.""" + srv = QueueServer(ToolCallEngine()) + out = _chat_once(srv, {"messages": [{"role": "user", "content": "hi"}]}) + ch = out["choices"][0] + assert "tool_calls" not in ch["message"] + assert "" in ch["message"]["content"] + assert ch["finish_reason"] == "stop" + + +def test_chat_once_parses_tool_calls_when_requested(): + srv = QueueServer(ToolCallEngine()) + out = _chat_once(srv, { + "messages": [{"role": "user", "content": "compute 2 + 2"}], + "tools": _TOOLS, "tool_choice": "auto", + }) + ch = out["choices"][0] + assert ch["message"]["content"] is None + assert ch["finish_reason"] == "tool_calls" + calls = ch["message"]["tool_calls"] + assert len(calls) == 1 + fn = calls[0]["function"] + assert fn["name"] == "calculator" + assert json.loads(fn["arguments"]) == {"expression": "2 + 2"} + assert calls[0]["type"] == "function" + assert calls[0]["id"].startswith("call_") + + +def test_chat_once_tool_choice_none_suppresses_parsing(): + srv = QueueServer(ToolCallEngine()) + out = _chat_once(srv, { + "messages": [{"role": "user", "content": "compute 2 + 2"}], + "tools": _TOOLS, "tool_choice": "none", + }) + ch = out["choices"][0] + assert "tool_calls" not in ch["message"] + assert "" in ch["message"]["content"] + + +def test_chat_once_tools_requested_but_engine_cannot_parse(): + """An engine with no parse_tool_calls (unsupported family) must not + crash a tools-enabled request; it just can't extract structured + calls, same as today.""" + srv = QueueServer(FakeEngine()) + out = _chat_once(srv, { + "messages": [{"role": "user", "content": "hi"}], "tools": _TOOLS, + }) + ch = out["choices"][0] + assert "tool_calls" not in ch["message"] + assert ch["finish_reason"] == "stop" + + +def test_chat_stream_tools_requested_buffers_and_emits_final_tool_calls(): + srv = QueueServer(ToolCallEngine()) + events = _chat_stream(srv, { + "messages": [{"role": "user", "content": "compute 2 + 2"}], + "stream": True, "tools": _TOOLS, "tool_choice": "auto", + }) + body = b"".join(events) + frames = [e for e in body.decode("utf-8").split("\n\n") if e] + assert frames[-1] == "data: [DONE]" + chunks = [json.loads(e[6:]) for e in frames[:-1]] + # No raw XML in any per-token delta -- buffered, not + # streamed token-by-token, unlike the no-tools path. + for c in chunks[:-1]: + assert "" not in json.dumps(c) + final = chunks[-1]["choices"][0] + assert final["finish_reason"] == "tool_calls" + assert final["delta"]["content"] is None + calls = final["delta"]["tool_calls"] + assert calls[0]["index"] == 0 + assert calls[0]["function"]["name"] == "calculator" + assert json.loads(calls[0]["function"]["arguments"]) == { + "expression": "2 + 2"} + + +def test_chat_stream_indexes_multiple_tool_calls(): + class MultiCallTok(ToolCallTok): + def decode(self, tokens): + return super().decode(tokens) * 2 + + server = QueueServer(ToolCallEngine(tok=MultiCallTok())) + frames = b"".join(_chat_stream(server, { + "messages": [{"role": "user", "content": "hi"}], "tools": _TOOLS, + })).decode().split("\n\n") + calls = json.loads(frames[-3][6:])["choices"][0]["delta"]["tool_calls"] + assert [call["index"] for call in calls] == [0, 1] + assert calls[0]["id"] != calls[1]["id"] + + +def test_chat_stream_without_tools_keeps_immediate_per_token_deltas(): + """Regression guard: requests with no tools must keep the existing + immediate per-token streaming -- buffering is scoped to tool-enabled + requests only.""" + srv = QueueServer(ToolCallEngine()) + events = _chat_stream(srv, { + "messages": [{"role": "user", "content": "hi"}], "stream": True, + }) + body = b"".join(events) + frames = [e for e in body.decode("utf-8").split("\n\n") if e] + chunks = [json.loads(e[6:]) for e in frames[:-1]] + deltas = [c["choices"][0]["delta"].get("content") + for c in chunks if c["choices"][0]["delta"]] + # unchanged from the no-tools path: whatever decode_tokens([tid]) + # returns per call, not the whole buffered response + assert len(deltas) == 3 + assert chunks[-1]["choices"][0]["finish_reason"] == "stop" + + def test_chat_stream_events_sequence(): srv = QueueServer(FakeEngine()) events = _chat_stream(srv, { @@ -353,6 +653,47 @@ def test_chat_stream_yields_before_generation_finishes(): # ---- stdlib HTTP transport ------------------------------------------------ +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize("transport", ["stdlib", "flask"]) +def test_http_rejects_invalid_tool_arguments_before_generation(stream, transport): + engine = FakeEngine() + server = QueueServer(engine) + payload = {"stream": stream, "messages": [{ + "role": "assistant", "content": None, + "tool_calls": [{"type": "function", "function": { + "name": "calculator", "arguments": "{broken"}}], + }]} + if transport == "flask": + if not _HAS_FLASK: + pytest.skip("flask is not installed") + response = create_app(server).test_client().post( + "/v1/chat/completions", json=payload) + assert response.status_code == 400 + error = response.get_json() + else: + handler = type("Edge0Handler", (_StdlibHandler,), {"server_q": server}) + httpd = ThreadingHTTPServer(("127.0.0.1", 0), handler) + thread = threading.Thread(target=httpd.serve_forever, daemon=True) + thread.start() + try: + req = urlrequest.Request( + f"http://127.0.0.1:{httpd.server_address[1]}/v1/chat/completions", + data=json.dumps(payload).encode(), + headers={"Content-Type": "application/json"}, + ) + with pytest.raises(HTTPError) as exc: + urlrequest.urlopen(req, timeout=10) + with exc.value as response: + assert response.code == 400 + error = json.loads(response.read()) + finally: + httpd.shutdown() + httpd.server_close() + thread.join(timeout=5) + assert "arguments must be a JSON object" in error["error"]["message"] + assert engine.generated == [] + + def test_stdlib_http_end_to_end(): eng = FakeEngine() srv = QueueServer(eng, model_name="fake") diff --git a/python/tests/test_tool_calls.py b/python/tests/test_tool_calls.py new file mode 100644 index 0000000..1b41187 --- /dev/null +++ b/python/tests/test_tool_calls.py @@ -0,0 +1,164 @@ +"""edge0.server.tool_calls: XML tool-call parsing for both checkpoint +families. + +The fixture strings below are not invented: they were captured by +rendering the *real* ``chat_template.jinja`` shipped with each checkpoint +(edge0-8b / Ling, edge0-35b / Qwen3.5) with jinja2 directly, for an +assistant message carrying one ``tool_calls`` entry, and copying the +generated span verbatim -- including the Ling template's own leading- +whitespace quirk before the first ```` (a real jinja2 whitespace- +control artifact in that template, not a typo here). This guards against +inventing a plausible-looking XML dialect that the real templates don't +actually produce. +""" + +from __future__ import annotations + +import json + +import pytest + +from edge0.server.tool_calls import parse_ling_tool_calls, parse_qwen_tool_calls + +# Captured from edge0-8b/chat_template.jinja rendered with +# tools=[calculator] and one assistant tool_calls entry. +LING_REAL_SPAN = ( + "calculator\n" + " expression\n" + "2 + 2\n" + "" +) + +# Captured from edge0-35b/chat_template.jinja's own in-prompt format +# example (example_function_name / example_parameter_1/2, the second +# spanning multiple lines -- the format the "IMPORTANT" reminder block +# documents to the model). +QWEN_REAL_SPAN = ( + "\n" + "\n" + "\n" + "value_1\n" + "\n" + "\n" + "This is the value for the second parameter\n" + "that can span\n" + "multiple lines\n" + "\n" + "\n" + "" +) + + +def test_ling_parses_real_template_span(): + content, calls = parse_ling_tool_calls(LING_REAL_SPAN) + assert content is None + assert len(calls) == 1 + fn = calls[0]["function"] + assert fn["name"] == "calculator" + assert json.loads(fn["arguments"]) == {"expression": "2 + 2"} + assert calls[0]["type"] == "function" + assert calls[0]["id"].startswith("call_") + + +def test_qwen_parses_real_template_span_with_multiline_value(): + content, calls = parse_qwen_tool_calls(QWEN_REAL_SPAN) + assert content is None + assert len(calls) == 1 + fn = calls[0]["function"] + assert fn["name"] == "example_function_name" + args = json.loads(fn["arguments"]) + assert args["example_parameter_1"] == "value_1" + assert args["example_parameter_2"] == ( + "This is the value for the second parameter\n" + "that can span\nmultiple lines") + + +def test_no_tool_call_leaves_content_unchanged_stripped(): + for parser in (parse_ling_tool_calls, parse_qwen_tool_calls): + content, calls = parser(" just a plain answer ") + assert content == "just a plain answer" + assert calls == [] + + +def test_prose_before_tool_call_is_preserved_as_content(): + text = "Sure, let me compute that.\n" + LING_REAL_SPAN + content, calls = parse_ling_tool_calls(text) + assert content == "Sure, let me compute that." + assert len(calls) == 1 + + +def test_multiple_tool_calls_in_one_response(): + text = LING_REAL_SPAN + "\n" + ( + "calculator\nexpression\n" + "3 * 3\n") + content, calls = parse_ling_tool_calls(text) + assert content is None + assert len(calls) == 2 + assert json.loads(calls[0]["function"]["arguments"]) == { + "expression": "2 + 2"} + assert json.loads(calls[1]["function"]["arguments"]) == { + "expression": "3 * 3"} + # each call gets its own id + assert calls[0]["id"] != calls[1]["id"] + + +def test_qwen_malformed_block_kept_as_content(): + # missing '' -- not what the template ever emits, but a + # truncated generation could still produce it; must not crash or + # silently eat the block. + text = "\nnot a function block\n" + content, calls = parse_qwen_tool_calls(text) + assert calls == [] + assert "" in content + + +@pytest.mark.parametrize("parser,body", [ + (parse_ling_tool_calls, "calculatorx"), + (parse_ling_tool_calls, "calculator1"), + (parse_ling_tool_calls, "calculatorx1junk"), + (parse_qwen_tool_calls, "broken"), + (parse_qwen_tool_calls, "junk"), + (parse_qwen_tool_calls, "\n1\n"), +]) +def test_incomplete_or_unconsumed_arguments_are_preserved(parser, body): + text = f"{body}" + content, calls = parser(text) + assert content == text + assert calls == [] + + +@pytest.mark.parametrize("parser,text", [ + (parse_ling_tool_calls, "get_time"), + (parse_qwen_tool_calls, ""), +]) +def test_zero_argument_tool_call(parser, text): + content, calls = parser(text) + assert content is None + assert json.loads(calls[0]["function"]["arguments"]) == {} + + +def test_qwen_function_allows_surrounding_whitespace(): + text = "\n \n " + content, calls = parse_qwen_tool_calls(text) + assert content is None + assert calls[0]["function"]["name"] == "get_time" + + +def test_malformed_blocks_survive_between_valid_calls(): + malformed = "calculatorx" + content, calls = parse_ling_tool_calls(LING_REAL_SPAN + malformed + LING_REAL_SPAN) + assert content == malformed + assert len(calls) == 2 + + +def test_ling_argument_coercion_numbers_and_json(): + text = ( + "set_temperature" + "value\n21.5" + "enabled\ntrue" + "label\nliving room" + "\n" + ) + _, calls = parse_ling_tool_calls(text) + args = json.loads(calls[0]["function"]["arguments"]) + assert args == {"value": 21.5, "enabled": True, "label": "living room"}