""" Agentic orchestration loop for the VL model. Calls `InferenceClient.chat_completion` with the three agent tools defined in `agent/tools.py`. The loop continues as long as the model emits tool calls. It stops (and returns a final string) when: - The model calls `final_output` -> return the `answer` argument - The model calls `abort` -> return "Aborted: {reason}" - The model returns plain content with no tool calls - `max_tool_rounds` is reached -> return the last assistant content """ from __future__ import annotations import json from typing import Any from huggingface_hub import InferenceClient from .tools import AGENT_TOOLS, dispatch_tool DEFAULT_MAX_TOOL_ROUNDS = 10 # Qwen3 "Thinking" models emit reasoning between these tags; the user-facing # answer follows the closing tag. See vLLM Qwen3ReasoningParser. _QWEN_THINK_START = "" _QWEN_THINK_END = "" def strip_qwen_thinking(text: str) -> str: """Drop Qwen3-Thinking chain-of-thought; keep only the final answer segment.""" if not text or not isinstance(text, str): return text s = text before, found, after = s.partition(_QWEN_THINK_START) s = after if found else before # Use the *last* closing tag: some replies contain multiple `` # segments; only the final segment is meant for the user. if _QWEN_THINK_END in s: return s.rsplit(_QWEN_THINK_END, 1)[-1].lstrip() return s def run_agent( messages: list[dict[str, Any]], client: InferenceClient, model: str, max_tokens: int = 512, temperature: float = 0.7, top_p: float = 0.95, max_tool_rounds: int = DEFAULT_MAX_TOOL_ROUNDS, ) -> str: """ Run the agentic loop and return the final answer as a string. Parameters ---------- messages: Full conversation so far, including the system message, all history, and the latest user message (which may contain an image as a multimodal content list). client: An authenticated `InferenceClient` instance. model: HF model ID to use for inference. max_tokens: Maximum tokens per completion call. temperature: Sampling temperature. top_p: Nucleus sampling top-p. max_tool_rounds: Hard cap on how many tool-calling rounds are allowed before the loop gives up and returns whatever the model last said. """ messages = list(messages) # work on a local copy for _round in range(max_tool_rounds + 1): response = client.chat_completion( messages=messages, model=model, tools=AGENT_TOOLS, tool_choice="auto", max_tokens=max_tokens, temperature=temperature, top_p=top_p, stream=False, ) choice = response.choices[0] msg = choice.message tool_calls = getattr(msg, "tool_calls", None) or [] # ------------------------------------------------------------------ # # No tool calls → plain text answer, we're done # # ------------------------------------------------------------------ # if not tool_calls: return _visible_assistant_text(msg) # ------------------------------------------------------------------ # # There are tool calls → check for terminal tools first # # ------------------------------------------------------------------ # for tc in tool_calls: fn = tc.function name = fn.name raw_args = fn.arguments or "{}" args = raw_args if isinstance(raw_args, dict) else _safe_parse(raw_args) if name == "final_output": return strip_qwen_thinking( args.get("answer", "") or _visible_assistant_text(msg) ) if name == "abort": reason = args.get("reason", "") return f"Aborted: {reason}" if reason else "Task aborted." # ------------------------------------------------------------------ # # Non-terminal tool calls → execute each one and feed results back # # ------------------------------------------------------------------ # messages.append(_assistant_tool_call_message(msg)) for tc in tool_calls: fn = tc.function name = fn.name raw_args = fn.arguments or "{}" args = raw_args if isinstance(raw_args, dict) else _safe_parse(raw_args) result = dispatch_tool(name, args) messages.append( { "role": "tool", "content": result, "tool_call_id": tc.id, } ) # Max rounds exhausted — return the last assistant content if any last_assistant = next( (m["content"] for m in reversed(messages) if m.get("role") == "assistant"), "I was unable to complete the task within the allowed number of steps.", ) return strip_qwen_thinking( last_assistant or "I was unable to complete the task within the allowed number of steps." ) # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- def _visible_assistant_text(msg: Any) -> str: """ User-visible assistant string from `message.content` only. The Inference API may return chain-of-thought in a separate `reasoning` field; we never merge that into the UI. Strip any ``...`` blocks left in `content`. """ raw = getattr(msg, "content", None) or "" return strip_qwen_thinking(raw) def _safe_parse(raw: str) -> dict: """Parse a JSON string into a dict, returning {} on failure.""" try: return json.loads(raw) except (json.JSONDecodeError, TypeError): return {} def _assistant_tool_call_message(msg: Any) -> dict: """ Re-serialise the assistant message that contains tool_calls into the plain dict format expected when appended back to `messages`. """ tool_calls_payload = [] for tc in msg.tool_calls or []: fn = tc.function tool_calls_payload.append( { "id": tc.id, "type": "function", "function": { "name": fn.name, "arguments": ( fn.arguments if isinstance(fn.arguments, str) else json.dumps(fn.arguments) ), }, } ) return { "role": "assistant", # Drop thinking from prior step so follow-up completions are not primed # to dump huge `` blocks into `content` on later turns. "content": strip_qwen_thinking(msg.content or ""), "tool_calls": tool_calls_payload, }