diff --git a/studio/backend/core/inference/llama_cpp.py b/studio/backend/core/inference/llama_cpp.py index 8054bd2a19..6769434fbc 100644 --- a/studio/backend/core/inference/llama_cpp.py +++ b/studio/backend/core/inference/llama_cpp.py @@ -1615,6 +1615,19 @@ class LlamaCppBackend: conversation = list(messages) url = f"{self.base_url}/v1/chat/completions" + # Allow-list of tool names the caller enabled for this request. + # Structured tool_calls are already constrained by llama-server, + # but healed calls (parsed from raw text below) are + # not, so they must be checked against this set before execution -- + # otherwise a model can invoke a tool that was never offered + # (e.g. emitting as text when only web_search + # was enabled). + _allowed_tool_names = { + t.get("function", {}).get("name") + for t in (tools or []) + if t.get("function", {}).get("name") + } + for iteration in range(max_tool_iterations): if cancel_event is not None and cancel_event.is_set(): return @@ -1672,43 +1685,65 @@ class LlamaCppBackend: ): tool_calls = self._parse_tool_calls_from_text(content_text) if tool_calls: - # Strip the tool call markup from content. - # Use greedy match within blocks since they - # can contain arbitrary content including code. - import re - - # Strip ... blocks (greedy inside) - content_text = re.sub( - r".*?", - "", - content_text, - flags = re.DOTALL, - ) - # Strip unterminated ... to end - content_text = re.sub( - r".*$", - "", - content_text, - flags = re.DOTALL, - ) - # Strip bare ... blocks - content_text = re.sub( - r".*?", - "", - content_text, - flags = re.DOTALL, - ) - # Strip unterminated bare to end - content_text = re.sub( - r".*$", - "", - content_text, - flags = re.DOTALL, - ).strip() logger.info( f"Parsed {len(tool_calls)} tool call(s) from content text" ) + # Always strip tool-call XML from content_text when any tool + # calls are present. llama-server may return structured + # tool_calls AND also leave XML in the content + # field, which would leak into the chat UI and conversation. + if ( + auto_heal_tool_calls + and tool_calls + and ("" in content_text or ".*?", + "", + content_text, + flags = re.DOTALL, + ) + content_text = re.sub( + r".*$", + "", + content_text, + flags = re.DOTALL, + ) + content_text = re.sub( + r".*?", + "", + content_text, + flags = re.DOTALL, + ) + content_text = re.sub( + r".*$", + "", + content_text, + flags = re.DOTALL, + ).strip() + + # Reject any tool calls whose name was not enabled for this + # request. Healed calls come from arbitrary model text and + # would otherwise bypass the caller's tool allow-list. + if tool_calls: + _kept = [ + tc + for tc in tool_calls + if tc.get("function", {}).get("name") in _allowed_tool_names + ] + if len(_kept) != len(tool_calls): + _dropped = [ + tc.get("function", {}).get("name") for tc in tool_calls + ] + logger.warning( + "Dropped tool call(s) not in enabled tools " + f"{sorted(_allowed_tool_names)}: {_dropped}" + ) + tool_calls = _kept + if finish_reason == "tool_calls" or (tool_calls and len(tool_calls) > 0): # Append the assistant message with tool_calls to conversation assistant_msg = {"role": "assistant", "content": content_text} @@ -1734,6 +1769,14 @@ class LlamaCppBackend: else: arguments = raw_args + # Malformed tool calls can carry non-object arguments + # (e.g. a JSON array or number parsed from a + # {...} payload); normalize to a + # dict so the .get() lookups below cannot raise + # AttributeError and abort the response. + if not isinstance(arguments, dict): + arguments = {} + # Yield status update if tool_name == "web_search": status_text = f"Searching: {arguments.get('query', '')}" @@ -1815,7 +1858,11 @@ class LlamaCppBackend: # Clear status yield {"type": "status", "text": ""} - # Final streaming pass with the full conversation context + # Final streaming pass with the full conversation context. + # Add stop sequences so the model cannot emit tool-call XML -- + # the non-streaming loop above already handled all tool + # iterations. If the model tries to call tools here it will + # simply stop, and we yield whatever text came before. stream_payload = { "messages": conversation, "stream": True, @@ -1832,19 +1879,16 @@ class LlamaCppBackend: } if max_tokens is not None: stream_payload["max_tokens"] = max_tokens - if stop: - stream_payload["stop"] = stop + _stop = list(stop) if stop else [] + if auto_heal_tool_calls: + _stop += ["", ".*?", _re_final.DOTALL), _re_final.compile(r".*?", _re_final.DOTALL), - ] - # Open-ended patterns strip from an opening tag to end-of-string. - # Only applied on the final flush to avoid non-monotonic shrinking. - _TOOL_ALL_PATTERNS = _TOOL_CLOSED_PATTERNS + [ _re_final.compile(r".*$", _re_final.DOTALL), _re_final.compile(r".*$", _re_final.DOTALL), ] @@ -1852,8 +1896,7 @@ class LlamaCppBackend: def _strip_tool_markup(text: str, *, final: bool = False) -> str: if not auto_heal_tool_calls: return text - patterns = _TOOL_ALL_PATTERNS if final else _TOOL_CLOSED_PATTERNS - for pat in patterns: + for pat in _TOOL_PATTERNS: text = pat.sub("", text) return text.strip() if final else text