Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 17 additions & 2 deletions unsloth_zoo/diffusion_studio/shim.py
Original file line number Diff line number Diff line change
Expand Up @@ -217,10 +217,25 @@ def health():
return {"status": "ok", "model": MODEL_ID}


def _tools_for_choice(tools, tool_choice):
# Honor tool_choice before advertising: "none" hides tools, a forced function
# narrows to just that one, anything else forwards all.
if isinstance(tool_choice, str) and tool_choice.lower() == "none":
return None
if isinstance(tool_choice, dict):
name = (tool_choice.get("function") or {}).get("name")
if name:
return [t for t in tools or []
if isinstance(t, dict) and (t.get("function") or {}).get("name") == name] or None
return tools
Comment on lines +228 to +230

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Preserve forced tool_choice semantics

When a client passes a forced tool_choice for one function, this only narrows the advertised schemas and then discards the actual choice before calling the visual server. The downstream request is therefore indistinguishable from auto with one tool, so the model can still return plain text instead of the required function call; OpenAI-compatible clients that force a tool rely on that call being enforced or rejected, not treated as optional.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Same as the earlier forced-tool_choice item: the diffusion visual server samples with no grammar (data.grammar is never applied) and never reads tool_choice, so required/forced enforcement is not achievable on this path regardless of forwarding. _tools_for_choice still narrows to the forced function for advertising.



@app.post("/v1/chat/completions")
async def chat(req: Request):
body = await req.json()
messages = body.get("messages", [])
# forwarded to the visual server, honoring tool_choice
tools = _tools_for_choice(body.get("tools"), body.get("tool_choice"))

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Normalize optional tool schema fields

When an OpenAI-compatible client sends a valid function tool that only includes the required name and omits description or parameters, this path forwards the raw tool to the visual server, whose common_chat_tools_parse_oaicompat reads those fields unconditionally. Those requests now fail with ERR parse (or stream an engine error) instead of completing; please add defaults for omitted optional fields before forwarding or reject them with a clean 400 response.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

common_chat_tools_parse_oaicompat uses function.value("description", "") and function.value("parameters", {}) at chat.cpp:551-555; only name uses .at(). A name-only tool parses without error, so there is no ERR parse to guard against.

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Emit tool calls before advertising tools

When an OpenAI-compatible request includes tools (especially with a forced tool_choice), this now advertises those schemas to the visual server, but the shim still returns the committed generation only as plain content and always finishes with "stop"; there is no conversion of generated native tool-call text into message.tool_calls or streamed delta.tool_calls. Tool-aware clients will therefore not execute the selected tool and may display the raw call instead, so the shim needs to parse and emit the Chat Completions tool-call fields before forwarding tool schemas.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Same as the earlier items on this file: this shim is an internal loopback bridge that Studio spawns, not a public entry point. The Studio backend parses the native tool-call text into structured tool_calls (unslothai/unsloth#6851), so emitting them here as well would double-process the same generation. Forwarding the schemas is what makes the visual server render them into the chat template; the conversion to Chat Completions tool-call fields happens one layer up, not in the shim.

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Forward the required tool choice state

When a client sends a forced tool choice (for example {"type":"function","function":{"name":"foo"}}) or tool_choice: "required", this line collapses the request to just a filtered tools list; _send then serializes only tools, so the child cannot distinguish an optional single tool from a required/forced call. Agent clients that rely on forced tool use can receive a normal assistant message instead of a tool call even though they explicitly required one.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Same as the earlier forced-tool_choice item: the diffusion visual server samples with no grammar and never reads tool_choice, so required/forced enforcement is not achievable on this path even if the field were forwarded. _tools_for_choice already narrows to the forced function so only that one schema is advertised, which is the most the server can act on here.

stream = bool(body.get("stream", False))
max_blocks = _max_blocks(body)
seed = int(body.get("seed", 3407))
Expand All @@ -234,7 +249,7 @@ async def chat(req: Request):
def work():
with _LOCK:
return V.generate_visual(srv, messages, seed=seed, max_blocks=max_blocks,
on_stats=stats_box.update)
on_stats=stats_box.update, tools=tools)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Return parsed tool calls from the shim endpoint

For direct /v1/chat/completions callers that send tools, this forwarding can make the visual server generate Gemma's native <|tool_call>... text, but the shim still returns the generated text only as assistant content with finish_reason: "stop" (and streams it as content deltas on the streaming path). OpenAI-compatible clients hitting this endpoint directly will not see message.tool_calls / finish_reason: "tool_calls", so they cannot execute the requested tool; please parse the native tool-call text here before returning or avoid enabling tool rendering in this shim until that response healing is local.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

By design: this shim is an internal loopback bridge that Studio spawns, and the Studio backend heals the native tool-call text into structured tool_calls (unslothai/unsloth#6851). On the supported Studio path the call is surfaced as tool_calls; a direct client on the loopback port is not a supported entry point, and emitting tool_calls in the shim would double-process against that healing.

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Return structured tool calls from the shim

When tools are supplied, this new argument lets the visual server render schemas and the model can respond with a tool-call marker, but chat() still funnels the returned text through _split_thought_channels and emits it only as message.content/content deltas with finish_reason: "stop" in the non-streaming and streaming branches. Non-Studio OpenAI clients will never see message.tool_calls/finish_reason: "tool_calls", so they won't execute the function even though the request advertised tools.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Same as the earlier item: this shim is an internal loopback bridge Studio spawns, and the Studio backend heals the native tool-call text into structured tool_calls (unslothai/unsloth#6851). A direct client on the loopback port is not a supported entry point, and emitting tool_calls here would double-process against that healing.

try:
text = await loop.run_in_executor(None, work)
except V.ContextOverflow as exc:
Expand Down Expand Up @@ -276,7 +291,7 @@ def work():
with _LOCK:
full = V.generate_visual(srv, messages, seed=seed, max_blocks=max_blocks,
on_frame=on_frame, on_commit=on_commit,
on_stats=stats_box.update)
on_stats=stats_box.update, tools=tools)
loop.call_soon_threadsafe(q.put_nowait, ("done", full))
except V.ContextOverflow as exc: # context budget exceeded -> clean user-facing message
loop.call_soon_threadsafe(q.put_nowait, ("overflow", exc))
Expand Down
20 changes: 12 additions & 8 deletions unsloth_zoo/diffusion_studio/visual_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -216,14 +216,18 @@ def restart(self):
pass
self._spawn()

def _send(self, messages, n_blocks, seed):
def _send(self, messages, n_blocks, seed, tools=None):
# A previous turn may have crashed the decoder; respawn before writing so a dead child does not
# strand this and every later turn with a broken pipe.
if self.p is None or self.p.poll() is not None:
self.restart()
with open(self.req, "w") as f:
json.dump({"seed": int(seed), "n_blocks": int(n_blocks), "messages": messages}, f,
ensure_ascii=False)
req = {"seed": int(seed), "n_blocks": int(n_blocks), "messages": messages}
# Forward tools so the server renders them into the chat template (inputs.tools)
# for schema-correct <|tool_call> args.
if tools:
req["tools"] = tools
Comment thread
oobabooga marked this conversation as resolved.
Comment on lines +227 to +228

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Preserve forced tool-choice semantics

When a request uses tool_choice: "required" or forces a specific function, the shim may narrow the registry, but the request file built here still only contains tools and drops the choice itself before the visual server renders the chat template. In those scenarios the server sees optional tools, so an OpenAI-compatible caller that explicitly required a tool call can still receive ordinary assistant content. Please forward a required/forced choice alongside the filtered registry, or reject unsupported choices, instead of only writing tools.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The diffusion visual server renders only the prompt and samples with no grammar (data.grammar is never applied) and never reads tool_choice, so required/forced enforcement is not achievable on this path regardless of forwarding. _tools_for_choice still narrows to the forced function for advertising, so forwarding the choice would be a no-op here.

with open(self.req, "w", encoding="utf-8") as f:
json.dump(req, f, ensure_ascii=False)
Comment on lines +228 to +230

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Write tool request JSON as UTF-8

On Windows or any non-UTF-8 locale, this new forwarding path serializes tool schemas through json.dump(..., ensure_ascii=False) into a file opened with the platform default encoding. An otherwise ASCII chat can now fail before reaching the visual server when a tool description contains characters such as °C or emoji, because those tool strings were not written at all before this change; open the request file with UTF-8 or escape non-ASCII when dumping the request.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed in 741b83f: the request file is now opened with encoding="utf-8", so non-ASCII tool descriptions and message content no longer depend on the platform default codec (this also covered the pre-existing message path). Confirmed the same content raises UnicodeEncodeError under cp1252.

try:
self.p.stdin.write(self.req + "\n")
self.p.stdin.flush()
Expand Down Expand Up @@ -256,7 +260,7 @@ def _parse_stats(line):
return stats


def generate_visual(server, messages, seed=3407, max_blocks=8, on_frame=None, on_commit=None, on_stats=None):
def generate_visual(server, messages, seed=3407, max_blocks=8, on_frame=None, on_commit=None, on_stats=None, tools=None):
"""Stream one turn through the optimized visual server.

on_frame(block, step, total, text): a denoising frame (the current argmax canvas, already decoded).
Expand All @@ -282,16 +286,16 @@ def _commit(txt):

for attempt in (0, 1):
try:
return _generate_visual_once(server, messages, seed, max_blocks, _frame, _commit, on_stats)
return _generate_visual_once(server, messages, seed, max_blocks, _frame, _commit, on_stats, tools)
except VisualServerCrashed:
server.restart() # always bring a fresh server up so the next turn works regardless
if attempt == 1 or progressed["emitted"]:
raise
# nothing emitted yet -> safe to transparently resend on the fresh server


def _generate_visual_once(server, messages, seed, max_blocks, on_frame, on_commit, on_stats):
server._send(messages, max_blocks, seed)
def _generate_visual_once(server, messages, seed, max_blocks, on_frame, on_commit, on_stats, tools=None):
server._send(messages, max_blocks, seed, tools)

full_text = ""
while True:
Expand Down
Loading