File size: 11,927 Bytes
d6a179b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
"""Pure parser for the gemma-4 coder's native/leaked tool markup β€” no dependencies.

The gemma-4 coder GGUF emits tool calls as text in `content` (`<|tool_call>call:NAME{...}`,
`<|tool>NAME{...}`, `<call:NAME{...}>`, `NAME(k=v)<|/tool|>`) instead of structured
tool_calls, because `llama.cpp --jinja` doesn't recognise its native format. This module
is the single source of truth for detecting those calls and parsing their arguments,
plus the serve-side tools-in-prompt block. Wrap it in a litellm callback (see
litellm_shim.py) or call it directly (see standalone.py). Pure stdlib, unit-tested.
"""
from __future__ import annotations

import json
import re
from typing import Any

# Stop strings: the turn separators the broken GGUF never treats as EOG β€” they bound
# the runaway `<|turn>user` loop. Injected by both the shim (pre_call) and the steer
# sweep (so its probes get a clean, bounded response instead of runaway garbage). We
# deliberately do NOT stop on `<channel|>` (the model may emit a thinking channel
# before the call); the parser takes the first call and ignores trailing hallucination.
STOP_STRINGS = ["<|turn>", "<turn|>", "<|turn|>"]

# Fancy/single quotes the model emits inside arg strings -> plain double quote.
_FANCY = str.maketrans({"’": '"', "β€˜": '"', "β€œ": '"', "”": '"', "'": '"'})
# Leaked tool-call markers vary by build / token-typing damage:
#   b9755:  <|tool>NAME{...}            (or <|tool|> / <tool|>)
#   b8733:  <call:NAME{...}>            (the <|tool_call> token stripped to '<')
#   raw:    <|tool_call>call:NAME{...}
# Anchor on a tool marker OR a `call:` prefix, then NAME{ body }.
# NAME char class = MCP's tool-name set (SEP-986): [A-Za-z0-9_.-], which is OpenAI's /
# Anthropic's ^[a-zA-Z0-9_-]{1,64}$ PLUS the dot used for namespacing (weather.get_weather,
# fs.read_file). The dot is the one that bit v3: a bare `[A-Za-z_]\w*` truncated
# `weather.get_weather`β†’`weather`, then no `{` β‡’ the call was dropped. First char is held
# to [A-Za-z_] (stricter than the specs, which permit a leading digit) as parser safety β€”
# real names don't start with a digit and it avoids matching `call:123{...}` junk.
_TOOL_RE = re.compile(
    r"(?:<\|?tool(?:_call)?\|?>\s*(?:call:)?|call:)\s*([A-Za-z_][\w.\-]*)\s*\{(.*?)\}",
    re.S,
)
# Paren form the model also leaks: NAME(k=v, ...) instead of NAME{...}. This is a
# *coder* model, so a bare NAME(...) is almost always real code β€” match ONLY when
# anchored on a tool marker: a LEADING <|tool>/<|tool_call> OR a TRAILING tool-close
# token (<|/tool|>, <tool|>, …). Body is paren-balanced-free ([^()]*).
_TOOL_PAREN_LEAD = re.compile(
    r"<\|?tool(?:_call)?\|?>\s*(?:call:)?\s*([A-Za-z_][\w.\-]*)\s*\(([^()]*)\)",
    re.S,
)
_TOOL_PAREN_CLOSE = re.compile(
    r"\b([A-Za-z_][\w.\-]*)\s*\(([^()]*)\)\s*<\|?/?\s*tool(?:_call)?\|?>",
    re.S,
)
# strip a leaked thinking channel:  <|channel>...<channel|>
_CHANNEL_RE = re.compile(r"<\|?(?:channel|think|thought)\|?>.*?<(?:channel|think|thought)\|>", re.S)
# residual control-token junk to scrub. Matches ONLY known gemma control tokens
# (<|name>, <name|>, <|name|>) β€” never bare <>/brackets, so code output from this
# *coder* model (C++ generics, HTML, comparisons) is left intact.
_CTRL_NAMES = "tool_call|tool_response|tool|channel|turn|think|thought|start_of_turn|end_of_turn"
_JUNK_RE = re.compile(rf"<\|?(?:{_CTRL_NAMES})\|?>")


def _coerce(v: str) -> Any:
    v = v.strip()
    if len(v) >= 2 and v[0] in "\"'β€™β€˜β€œβ€" and v[-1] in "\"'β€™β€˜β€œβ€":
        return v[1:-1]
    low = v.lower()
    if low == "true":
        return True
    if low == "false":
        return False
    if low == "null":
        return None
    try:
        return int(v)
    except ValueError:
        try:
            return float(v)
        except ValueError:
            return v


def _split_top(s: str) -> list[str]:
    """Split on top-level commas only (ignore commas nested in {}/[]/())."""
    parts, depth, cur = [], 0, ""
    for ch in s:
        if ch in "{[(":
            depth += 1
        elif ch in "}])":
            depth = max(0, depth - 1)
        if ch == "," and depth == 0:
            parts.append(cur)
            cur = ""
        else:
            cur += ch
    if cur.strip():
        parts.append(cur)
    return parts


def _parse_args(body: str) -> dict:
    if not body.strip():
        return {}
    # Fast path: normalise fancy quotes, quote bare keys, try strict JSON.
    norm = body.translate(_FANCY)
    norm = re.sub(r"([{,]\s*)([A-Za-z_]\w*)\s*:", r'\1"\2":', "{" + norm + "}")
    try:
        out = json.loads(norm)
        if isinstance(out, dict):
            return out
    except Exception:
        pass
    # Tolerant fallback: split top-level commas, coerce scalars. Accept BOTH `:`
    # (brace form {k: v}) and `=` (paren form (k=v)) as the key/value separator.
    out: dict = {}
    for part in _split_top(body):
        m = re.search(r"[:=]", part)
        if not m:
            continue
        k, v = part[: m.start()], part[m.start() + 1 :]
        out[k.strip().strip("\"'β€™β€˜β€œβ€")] = _coerce(v)
    return out


def find_tool_calls(text: str) -> tuple[list[dict], str]:
    """Return (calls, leftover_content) where calls is a list of
    {"name": str, "arguments": dict} parsed from leaked gemma tool markup. Pure β€”
    no litellm types. The shim wraps this into OpenAI tool_calls; the sweep counts it."""
    if not text or ("tool" not in text and "call:" not in text):
        return [], text or ""
    # Gather candidates from the brace form {…} and both anchored paren forms (k=v),
    # then merge by position and drop overlaps (a call matching both a leading marker
    # and a trailing close shows up twice β€” keep one).
    raw: list[tuple[int, int, str, str]] = []
    for rx in (_TOOL_RE, _TOOL_PAREN_LEAD, _TOOL_PAREN_CLOSE):
        for m in rx.finditer(text):
            raw.append((m.start(), m.end(), m.group(1), m.group(2)))
    if not raw:
        return [], text
    raw.sort(key=lambda t: (t[0], -(t[1] - t[0])))  # earliest first; longer span wins a tie
    spans: list[tuple[int, int, str, str]] = []
    for s, e, name, body in raw:
        if any(s < pe and e > ps for ps, pe, _, _ in spans):  # overlaps a kept span
            continue
        spans.append((s, e, name, body))
    spans.sort(key=lambda t: t[0])
    calls = [{"name": name, "arguments": _parse_args(body)} for _, _, name, body in spans]
    first_start = spans[0][0]
    # Residual assistant content = only the preamble BEFORE the first call (everything
    # after is more calls or a hallucinated result). Scrub control tokens / thinking
    # channels; edge-strip the stray '<' a '<call:' wrapper leaves (edges only).
    preamble = _CHANNEL_RE.sub("", text[:first_start])
    leftover = _JUNK_RE.sub("", preamble).strip().strip("<>").strip()
    if len(re.sub(r"\W", "", leftover)) < 2:
        leftover = ""
    return calls, leftover


def repair_native_args(name: str, arguments: Any) -> tuple[str, bool]:
    """Repair the `arguments` string of a tool_call the llama.cpp NATIVE parser
    already captured (the canonical-call path: with --jinja a model whose calls
    the native peg-gemma4 parser recognises lands them in tool_calls.arguments,
    often as malformed JSON, with empty content β€” so the content shim never sees
    them). Returns (json_string, changed).

    Parity guarantee: if `arguments` already parses to a JSON object it is
    returned verbatim with changed=False. The production huihui coder leaks its
    calls into *content* and leaves native tool_calls empty, so this never fires
    for it β€” only a present-but-malformed native arg string is re-parsed, through
    the same tolerant parser the content path already uses.
    """
    if not isinstance(arguments, str):
        return arguments, False
    s = arguments.strip()
    try:
        if isinstance(json.loads(s), dict):
            return arguments, False  # already valid JSON object β€” leave verbatim
    except Exception:
        pass
    # Malformed: reconstruct a `call:NAME{body}` shape and reuse find_tool_calls.
    body = s if (s.startswith("{") and s.endswith("}")) else "{" + s + "}"
    calls, _ = find_tool_calls(f"call:{name}{body}")
    if calls:
        return json.dumps(calls[0]["arguments"], ensure_ascii=False), True
    # Last resort: parse the bare body directly (handles non-brace blobs).
    parsed = _parse_args(s.strip().lstrip("{").rstrip("}"))
    if parsed:
        return json.dumps(parsed, ensure_ascii=False), True
    return arguments, False


def clean_content(text: str) -> str:
    """Scrub the broken GGUF's leaked control tokens / thinking channels from plain
    content, without touching bare <>/brackets (code-safe)."""
    if not text:
        return text
    out = _CHANNEL_RE.sub("", text)
    out = _JUNK_RE.sub("", out)
    return out.strip()


# ── serve-side tools-in-prompt ─────────────────────────────────────────────────
# A model fine-tuned with tool DEFINITIONS folded into the prompt only behaves if it
# sees the SAME block at serve time. render_tools_prompt() reproduces that block; call
# fold_tools_prompt(messages, tools) before sending so train == serve.
TOOLS_HEADER = "# Available tools"
_NAME_BAD = re.compile(r"[^A-Za-z0-9_.\-]+")


def _sanitize_tool_name(name: str) -> str:
    clean = _NAME_BAD.sub("_", str(name)).strip("_") or "tool"
    if not re.match(r"[A-Za-z_]", clean):
        clean = "_" + clean
    return clean


def _normalize_tool_schema(tool: Any) -> dict | None:
    """OpenAI-nested or flat tool schema β†’ {name, args, required, description}, or None."""
    if isinstance(tool, dict) and isinstance(tool.get("function"), dict):
        tool = tool["function"]
    if not isinstance(tool, dict) or not tool.get("name"):
        return None
    params = tool.get("parameters")
    if not isinstance(params, dict):
        params = tool.get("arguments") if isinstance(tool.get("arguments"), dict) else {}
    props = params.get("properties")
    args = sorted(props) if isinstance(props, dict) else []
    required = params.get("required")
    required = [str(r) for r in required] if isinstance(required, list) else []
    return {
        "name": _sanitize_tool_name(tool["name"]),
        "args": args,
        "required": required,
        "description": str(tool.get("description") or "").strip(),
    }


def render_tools_prompt(tools: list | None) -> str:
    """Render tool DEFINITIONS as the prompt block, or "" if none valid. Mirrors the
    corpus-side render_tools_block (contract-tested for byte equality)."""
    norm = [t for t in (_normalize_tool_schema(x) for x in (tools or [])) if t]
    if not norm:
        return ""
    lines = [
        TOOLS_HEADER,
        "To use a tool, emit <|tool_call>call:NAME followed by a compact-JSON argument "
        "object. Available functions (required args marked *):",
    ]
    for t in norm:
        sig = ", ".join((a + "*" if a in t["required"] else a) for a in t["args"])
        desc = f" β€” {t['description']}" if t["description"] else ""
        lines.append(f"- {t['name']}({sig}){desc}")
    return "\n".join(lines)


def fold_tools_prompt(messages: list, tools: list | None) -> list:
    """Prepend the rendered tools block to the FIRST user message (serve-side mirror of
    the corpus fold). Returns a NEW list; no-op when there are no valid tools."""
    block = render_tools_prompt(tools)
    if not block or not isinstance(messages, list):
        return messages
    out = [dict(m) if isinstance(m, dict) else m for m in messages]
    for m in out:
        if isinstance(m, dict) and m.get("role") == "user" and isinstance(m.get("content"), str):
            m["content"] = f"{block}\n\n{m['content']}".strip()
            break
    return out