diff --git a/studio/backend/core/inference/llama_cpp.py b/studio/backend/core/inference/llama_cpp.py
index 3f89cd3e5d..eac54c4a53 100644
--- a/studio/backend/core/inference/llama_cpp.py
+++ b/studio/backend/core/inference/llama_cpp.py
@@ -57,6 +57,7 @@ class LlamaCppBackend:
self._stdout_lines: list[str] = []
self._stdout_thread: Optional[threading.Thread] = None
self._cancel_event = threading.Event()
+ self._api_key: Optional[str] = None
self._kill_orphaned_servers()
atexit.register(self._cleanup)
@@ -938,6 +939,17 @@ class LlamaCppBackend:
cmd.extend(["--mmproj", mmproj_path])
logger.info(f"Using mmproj for vision: {mmproj_path}")
+ # Option C: add --api-key for direct client access when enabled
+ import os as _os
+ import secrets as _secrets
+
+ if _os.getenv("UNSLOTH_DIRECT_STREAM", "0") == "1":
+ self._api_key = _secrets.token_urlsafe(32)
+ cmd.extend(["--api-key", self._api_key])
+ logger.info("llama-server started with --api-key for direct streaming")
+ else:
+ self._api_key = None
+
logger.info(f"Starting llama-server: {' '.join(cmd)}")
# Set library paths so llama-server can find its shared libs and CUDA DLLs
@@ -1407,6 +1419,7 @@ class LlamaCppBackend:
url: str,
payload: dict,
cancel_event: Optional[threading.Event] = None,
+ headers: Optional[dict] = None,
):
"""Open an httpx streaming POST with cancel support.
@@ -1473,7 +1486,8 @@ class LlamaCppBackend:
pool = 10,
)
with client.stream(
- "POST", url, json = payload, timeout = prefill_timeout
+ "POST", url, json = payload, timeout = prefill_timeout,
+ headers = headers,
) as response:
_response_ref[0] = response
if cancel_event is not None and cancel_event.is_set():
@@ -1547,9 +1561,10 @@ class LlamaCppBackend:
# can finish. Cancel during streaming is handled by the
# watcher thread (closes the response on cancel_event).
stream_timeout = httpx.Timeout(connect = 10, read = 0.5, write = 10, pool = 10)
+ _auth_headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else None
with httpx.Client(timeout = stream_timeout) as client:
with self._stream_with_retry(
- client, url, payload, cancel_event
+ client, url, payload, cancel_event, headers = _auth_headers,
) as response:
if response.status_code != 200:
error_body = response.read().decode()
@@ -1681,14 +1696,40 @@ class LlamaCppBackend:
_accumulated_predicted_ms = 0.0
_accumulated_predicted_n = 0
+ # ── Shared patterns for stripping tool XML from streamed content ──
+ import re as _re_tool
+
+ _TOOL_CLOSED_PATTERNS = [
+ _re_tool.compile(r".*?", _re_tool.DOTALL),
+ _re_tool.compile(r".*?", _re_tool.DOTALL),
+ ]
+ _TOOL_ALL_PATTERNS = _TOOL_CLOSED_PATTERNS + [
+ _re_tool.compile(r".*$", _re_tool.DOTALL),
+ _re_tool.compile(r".*$", _re_tool.DOTALL),
+ ]
+
+ 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:
+ text = pat.sub("", text)
+ return text.strip() if final else text
+
+ # XML prefixes that signal a tool call in content
+ _TOOL_XML_SIGNALS = ("", "= 0 else 0,
@@ -1706,65 +1747,327 @@ class LlamaCppBackend:
payload["stop"] = stop
try:
- with httpx.Client(timeout = None) as client:
- resp = client.post(url, json = payload)
- if resp.status_code != 200:
- raise RuntimeError(
- f"llama-server returned {resp.status_code}: {resp.text}"
+ _auth_headers = (
+ {"Authorization": f"Bearer {self._api_key}"}
+ if self._api_key else None
+ )
+
+ # ── Speculative buffer state machine ──────────────────
+ # BUFFERING: accumulating content, checking for tool signals
+ # STREAMING: no tool detected, yielding tokens to caller
+ # DRAINING: tool signal found, silently consuming rest
+ _S_BUFFERING = 0
+ _S_STREAMING = 1
+ _S_DRAINING = 2
+
+ detect_state = _S_BUFFERING
+ content_buffer = "" # Raw content held during BUFFERING
+ content_accum = "" # All content tokens (for tool parsing)
+ reasoning_accum = ""
+ cumulative_display = "" # Cumulative text yielded (with )
+ in_thinking = False
+ has_content_tokens = False
+ tool_calls_acc = {} # Structured delta.tool_calls fragments
+ has_structured_tc = False
+ _iter_usage = None
+ _iter_timings = None
+ _stream_done = False
+ _last_emitted = ""
+
+ stream_timeout = httpx.Timeout(
+ connect = 10, read = 0.5, write = 10, pool = 10,
+ )
+ with httpx.Client(timeout = stream_timeout) as client:
+ with self._stream_with_retry(
+ client, url, payload, cancel_event,
+ headers = _auth_headers,
+ ) as response:
+ if response.status_code != 200:
+ error_body = response.read().decode()
+ raise RuntimeError(
+ f"llama-server returned {response.status_code}: "
+ f"{error_body}"
+ )
+
+ raw_buf = ""
+ for raw_chunk in self._iter_text_cancellable(
+ response, cancel_event,
+ ):
+ raw_buf += raw_chunk
+ while "\n" in raw_buf:
+ line, raw_buf = raw_buf.split("\n", 1)
+ line = line.strip()
+
+ if not line:
+ continue
+ if line == "data: [DONE]":
+ # Flush thinking state for STREAMING
+ if detect_state == _S_STREAMING and in_thinking:
+ if has_content_tokens:
+ cumulative_display += ""
+ yield {
+ "type": "content",
+ "text": _strip_tool_markup(
+ cumulative_display,
+ final = True,
+ ),
+ }
+ else:
+ cumulative_display = reasoning_accum
+ yield {
+ "type": "content",
+ "text": cumulative_display,
+ }
+ _stream_done = True
+ break # exit inner while
+ if not line.startswith("data: "):
+ continue
+
+ try:
+ chunk_data = json.loads(line[6:])
+ _ct = chunk_data.get("timings")
+ if _ct:
+ _iter_timings = _ct
+ _cu = chunk_data.get("usage")
+ if _cu:
+ _iter_usage = _cu
+
+ choices = chunk_data.get("choices", [])
+ if not choices:
+ continue
+
+ delta = choices[0].get("delta", {})
+
+ # ── Structured tool_calls ──
+ tc_deltas = delta.get("tool_calls")
+ if tc_deltas:
+ has_structured_tc = True
+ detect_state = _S_DRAINING
+ for tc_d in tc_deltas:
+ idx = tc_d.get("index", 0)
+ if idx not in tool_calls_acc:
+ tool_calls_acc[idx] = {
+ "id": tc_d.get(
+ "id", f"call_{idx}"
+ ),
+ "type": "function",
+ "function": {
+ "name": "",
+ "arguments": "",
+ },
+ }
+ func = tc_d.get("function", {})
+ if func.get("name"):
+ tool_calls_acc[idx][
+ "function"
+ ]["name"] += func["name"]
+ if func.get("arguments"):
+ tool_calls_acc[idx][
+ "function"
+ ]["arguments"] += func[
+ "arguments"
+ ]
+ continue
+
+ # ── Reasoning tokens (bypass buffer) ──
+ reasoning = delta.get(
+ "reasoning_content", ""
+ )
+ if reasoning:
+ reasoning_accum += reasoning
+ if detect_state != _S_DRAINING:
+ if not in_thinking:
+ cumulative_display += ""
+ in_thinking = True
+ cumulative_display += reasoning
+ yield {
+ "type": "content",
+ "text": cumulative_display,
+ }
+
+ # ── Content tokens ──
+ token = delta.get("content", "")
+ if token:
+ has_content_tokens = True
+ content_accum += token
+
+ if detect_state == _S_DRAINING:
+ pass # accumulate silently
+
+ elif detect_state == _S_STREAMING:
+ if in_thinking:
+ cumulative_display += ""
+ in_thinking = False
+ cumulative_display += token
+ cleaned = _strip_tool_markup(
+ cumulative_display,
+ )
+ if len(cleaned) > len(
+ _last_emitted
+ ):
+ _last_emitted = cleaned
+ yield {
+ "type": "content",
+ "text": cleaned,
+ }
+
+ elif detect_state == _S_BUFFERING:
+ content_buffer += token
+ stripped_buf = (
+ content_buffer.lstrip()
+ )
+ if not stripped_buf:
+ continue
+
+ # Check tool signal prefixes
+ is_prefix = False
+ is_match = False
+ for sig in _TOOL_XML_SIGNALS:
+ if stripped_buf.startswith(
+ sig
+ ):
+ is_match = True
+ break
+ if sig.startswith(
+ stripped_buf
+ ):
+ is_prefix = True
+ break
+
+ if is_match:
+ detect_state = _S_DRAINING
+ elif (
+ is_prefix
+ and len(stripped_buf)
+ < _MAX_BUFFER_CHARS
+ ):
+ pass # keep buffering
+ else:
+ # Not a tool -- flush buffer
+ detect_state = _S_STREAMING
+ if in_thinking:
+ cumulative_display += (
+ ""
+ )
+ in_thinking = False
+ cumulative_display += (
+ content_buffer
+ )
+ cleaned = _strip_tool_markup(
+ cumulative_display,
+ )
+ if len(cleaned) > len(
+ _last_emitted
+ ):
+ _last_emitted = cleaned
+ yield {
+ "type": "content",
+ "text": cleaned,
+ }
+
+ except json.JSONDecodeError:
+ logger.debug(
+ f"Skipping malformed SSE line: "
+ f"{line[:100]}"
+ )
+ if _stream_done:
+ break # exit outer for
+
+ # ── Resolve BUFFERING at stream end ──
+ if detect_state == _S_BUFFERING:
+ stripped_buf = content_buffer.lstrip()
+ if stripped_buf and auto_heal_tool_calls and any(
+ s in stripped_buf for s in _TOOL_XML_SIGNALS
+ ):
+ detect_state = _S_DRAINING
+ elif content_accum or reasoning_accum:
+ detect_state = _S_STREAMING
+ if content_buffer:
+ if in_thinking:
+ cumulative_display += ""
+ in_thinking = False
+ cumulative_display += content_buffer
+ yield {
+ "type": "content",
+ "text": _strip_tool_markup(
+ cumulative_display, final = True,
+ ),
+ }
+ else:
+ return
+
+ # ── STREAMING path: no tool call ──
+ if detect_state == _S_STREAMING:
+ # Safety net: check for XML tool signals in content
+ _safety_tc = None
+ if auto_heal_tool_calls and any(
+ s in content_accum for s in _TOOL_XML_SIGNALS
+ ):
+ _safety_tc = self._parse_tool_calls_from_text(
+ content_accum,
)
- data = resp.json()
- except httpx.ConnectError:
- raise RuntimeError("Lost connection to llama-server")
+ if not _safety_tc:
+ # Content was already streamed. Yield metadata.
+ yield {"type": "status", "text": ""}
+ _fu = _iter_usage or {}
+ _fc = _fu.get("completion_tokens", 0)
+ _fp = _fu.get("prompt_tokens", 0)
+ _tc = _fc + _accumulated_completion_tokens
+ if _iter_usage or _iter_timings:
+ _mt = (
+ dict(_iter_timings) if _iter_timings else {}
+ )
+ if (
+ _accumulated_predicted_ms
+ or _accumulated_predicted_n
+ ):
+ _mt["predicted_ms"] = (
+ _mt.get("predicted_ms", 0)
+ + _accumulated_predicted_ms
+ )
+ _tn = (
+ _mt.get("predicted_n", 0)
+ + _accumulated_predicted_n
+ )
+ _mt["predicted_n"] = _tn
+ _tms = _mt["predicted_ms"]
+ if _tms > 0:
+ _mt["predicted_per_second"] = (
+ _tn / (_tms / 1000.0)
+ )
+ yield {
+ "type": "metadata",
+ "usage": {
+ "prompt_tokens": _fp,
+ "completion_tokens": _tc,
+ "total_tokens": _fp + _tc,
+ },
+ "timings": _mt,
+ }
+ return
- choices = data.get("choices", [])
- if not choices:
- return
-
- choice = choices[0]
- finish_reason = choice.get("finish_reason", "")
- message = choice.get("message", {})
-
- # If model wants to call tools
- tool_calls = message.get("tool_calls")
-
- # Fallback: detect tool calls embedded as XML/text in content
- # Some models output XML instead of structured tool_calls,
- # or bare tags without wrapper.
- content_text = message.get("content", "") or ""
- if (
- auto_heal_tool_calls
- and not tool_calls
- and ("" in content_text or " blocks since they
- # can contain arbitrary content including code.
+ # Safety net caught tool XML -- treat as tool call
+ tool_calls = _safety_tc
+ content_text = content_accum
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".*$",
"",
@@ -1772,30 +2075,86 @@ class LlamaCppBackend:
flags = re.DOTALL,
).strip()
logger.info(
- f"Parsed {len(tool_calls)} tool call(s) from content text"
+ f"Safety net: parsed {len(tool_calls)} tool call(s) "
+ f"from streamed content"
)
+ else:
+ # ── DRAINING path: assemble tool_calls ──
+ tool_calls = None
+ content_text = content_accum
+ if has_structured_tc:
+ tool_calls = [
+ tool_calls_acc[i]
+ for i in sorted(tool_calls_acc)
+ ]
+ if not tool_calls and auto_heal_tool_calls and any(
+ s in content_accum for s in _TOOL_XML_SIGNALS
+ ):
+ tool_calls = self._parse_tool_calls_from_text(
+ content_accum,
+ )
+ if tool_calls and not has_structured_tc:
+ import re
+ 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,
+ )
+ content_text = re.sub(
+ r".*$",
+ "",
+ content_text,
+ flags = re.DOTALL,
+ ).strip()
+ if tool_calls:
+ logger.info(
+ f"Parsed {len(tool_calls)} tool call(s) from "
+ f"{'structured delta' if has_structured_tc else 'content text'}"
+ )
+ if not tool_calls:
+ # DRAINING but no tool calls (false positive)
+ yield {"type": "status", "text": ""}
+ if content_accum:
+ yield {"type": "content", "text": content_accum}
+ if _iter_usage or _iter_timings:
+ yield {
+ "type": "metadata",
+ "usage": _iter_usage,
+ "timings": _iter_timings,
+ }
+ return
- if finish_reason == "tool_calls" or (tool_calls and len(tool_calls) > 0):
- # Only accumulate metrics for responses that are actually used
- _accumulated_completion_tokens += data.get("usage", {}).get(
- "completion_tokens", 0
+ # ── Execute tool calls ──
+ _accumulated_completion_tokens += (
+ (_iter_usage or {}).get("completion_tokens", 0)
)
- _iter_timings = data.get("timings", {})
- _accumulated_predicted_ms += _iter_timings.get("predicted_ms", 0)
- _accumulated_predicted_n += _iter_timings.get("predicted_n", 0)
- # Append the assistant message with tool_calls to conversation
+ _it = _iter_timings or {}
+ _accumulated_predicted_ms += _it.get("predicted_ms", 0)
+ _accumulated_predicted_n += _it.get("predicted_n", 0)
+
assistant_msg = {"role": "assistant", "content": content_text}
if tool_calls:
assistant_msg["tool_calls"] = tool_calls
conversation.append(assistant_msg)
- # Execute each tool call
for tc in tool_calls or []:
func = tc.get("function", {})
tool_name = func.get("name", "")
raw_args = func.get("arguments", {})
- # Handle arguments as either string or dict
if isinstance(raw_args, str):
try:
arguments = json.loads(raw_args)
@@ -1807,12 +2166,13 @@ class LlamaCppBackend:
else:
arguments = raw_args
- # Yield status update
if tool_name == "web_search":
status_text = f"Searching: {arguments.get('query', '')}"
elif tool_name == "python":
preview = (
- (arguments.get("code") or "").strip().split("\n")[0][:60]
+ (arguments.get("code") or "")
+ .strip()
+ .split("\n")[0][:60]
)
status_text = (
f"Running Python: {preview}"
@@ -1830,7 +2190,6 @@ class LlamaCppBackend:
status_text = f"Calling: {tool_name}"
yield {"type": "status", "text": status_text}
- # Emit tool_start so the frontend can record inputs
yield {
"type": "tool_start",
"tool_name": tool_name,
@@ -1838,7 +2197,6 @@ class LlamaCppBackend:
"arguments": arguments,
}
- # Execute the tool
_effective_timeout = (
None if tool_call_timeout >= 9999 else tool_call_timeout
)
@@ -1850,7 +2208,6 @@ class LlamaCppBackend:
session_id = session_id,
)
- # Emit tool_end so the frontend can record outputs
yield {
"type": "tool_end",
"tool_name": tool_name,
@@ -1858,7 +2215,6 @@ class LlamaCppBackend:
"result": result,
}
- # Append tool result to conversation
tool_msg = {
"role": "tool",
"name": tool_name,
@@ -1872,26 +2228,13 @@ class LlamaCppBackend:
# Continue the loop to let model respond with context
continue
- # No tool calls -- model answered directly.
- # If no tools were executed at all, just yield the content
- # from this response instead of making a redundant second request.
- if iteration == 0 and content_text:
- yield {"type": "status", "text": ""}
- yield {"type": "content", "text": content_text}
- _direct_usage = data.get("usage")
- _direct_timings = data.get("timings")
- if _direct_usage or _direct_timings:
- yield {
- "type": "metadata",
- "usage": _direct_usage,
- "timings": _direct_timings,
- }
- return
+ except httpx.ConnectError:
+ raise RuntimeError("Lost connection to llama-server")
+ except Exception as e:
+ if cancel_event is not None and cancel_event.is_set():
+ return
+ raise
- # Tools were called in previous iterations; do a final
- # streaming pass so the model can synthesize a response
- # incorporating the tool results.
- break
# Clear status
yield {"type": "status", "text": ""}
@@ -1917,27 +2260,6 @@ class LlamaCppBackend:
stream_payload["stop"] = stop
stream_payload["stream_options"] = {"include_usage": True}
- import re as _re_final
-
- # Closed blocks only -- safe to strip mid-stream without shrinking later.
- _TOOL_CLOSED_PATTERNS = [
- _re_final.compile(r".*?", _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),
- ]
-
- 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:
- text = pat.sub("", text)
- return text.strip() if final else text
cumulative = ""
_last_emitted = ""
@@ -1950,9 +2272,10 @@ class LlamaCppBackend:
try:
stream_timeout = httpx.Timeout(connect = 10, read = 0.5, write = 10, pool = 10)
+ _auth_headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else None
with httpx.Client(timeout = stream_timeout) as client:
with self._stream_with_retry(
- client, url, stream_payload, cancel_event
+ client, url, stream_payload, cancel_event, headers = _auth_headers,
) as response:
if response.status_code != 200:
error_body = response.read().decode()
@@ -2078,7 +2401,8 @@ class LlamaCppBackend:
if not self.is_loaded:
return None
try:
- with httpx.Client(timeout = 10) as client:
+ _auth_headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else {}
+ with httpx.Client(timeout = 10, headers = _auth_headers) as client:
def _detok(tid: int) -> str:
r = client.post(
@@ -2196,7 +2520,8 @@ class LlamaCppBackend:
if need_ids:
payload["n_probs"] = 1
- with httpx.Client(timeout = httpx.Timeout(300, connect = 10)) as client:
+ _auth_headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else {}
+ with httpx.Client(timeout = httpx.Timeout(300, connect = 10), headers = _auth_headers) as client:
resp = client.post(f"{self.base_url}/completion", json = payload)
if resp.status_code != 200:
raise RuntimeError(