From 190124d4acd392b6ccaf46dd922f33930a75c7ea Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 27 Mar 2026 06:17:46 +0000 Subject: [PATCH] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- studio/backend/core/inference/llama_cpp.py | 162 ++++++++++----------- 1 file changed, 81 insertions(+), 81 deletions(-) diff --git a/studio/backend/core/inference/llama_cpp.py b/studio/backend/core/inference/llama_cpp.py index eac54c4a53..c4901d4df2 100644 --- a/studio/backend/core/inference/llama_cpp.py +++ b/studio/backend/core/inference/llama_cpp.py @@ -1486,7 +1486,10 @@ 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 @@ -1561,10 +1564,16 @@ 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 + _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, headers = _auth_headers, + client, + url, + payload, + cancel_event, + headers = _auth_headers, ) as response: if response.status_code != 200: error_body = response.read().decode() @@ -1749,7 +1758,8 @@ class LlamaCppBackend: try: _auth_headers = ( {"Authorization": f"Bearer {self._api_key}"} - if self._api_key else None + if self._api_key + else None ) # ── Speculative buffer state machine ────────────────── @@ -1758,16 +1768,16 @@ class LlamaCppBackend: # DRAINING: tool signal found, silently consuming rest _S_BUFFERING = 0 _S_STREAMING = 1 - _S_DRAINING = 2 + _S_DRAINING = 2 detect_state = _S_BUFFERING - content_buffer = "" # Raw content held during BUFFERING - content_accum = "" # All content tokens (for tool parsing) + content_buffer = "" # Raw content held during BUFFERING + content_accum = "" # All content tokens (for tool parsing) reasoning_accum = "" - cumulative_display = "" # Cumulative text yielded (with ) + cumulative_display = "" # Cumulative text yielded (with ) in_thinking = False has_content_tokens = False - tool_calls_acc = {} # Structured delta.tool_calls fragments + tool_calls_acc = {} # Structured delta.tool_calls fragments has_structured_tc = False _iter_usage = None _iter_timings = None @@ -1775,11 +1785,17 @@ class LlamaCppBackend: _last_emitted = "" stream_timeout = httpx.Timeout( - connect = 10, read = 0.5, write = 10, pool = 10, + 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, + client, + url, + payload, + cancel_event, headers = _auth_headers, ) as response: if response.status_code != 200: @@ -1791,7 +1807,8 @@ class LlamaCppBackend: raw_buf = "" for raw_chunk in self._iter_text_cancellable( - response, cancel_event, + response, + cancel_event, ): raw_buf += raw_chunk while "\n" in raw_buf: @@ -1847,9 +1864,7 @@ class LlamaCppBackend: 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}" - ), + "id": tc_d.get("id", f"call_{idx}"), "type": "function", "function": { "name": "", @@ -1858,21 +1873,17 @@ class LlamaCppBackend: } func = tc_d.get("function", {}) if func.get("name"): - tool_calls_acc[idx][ - "function" - ]["name"] += func["name"] + tool_calls_acc[idx]["function"][ + "name" + ] += func["name"] if func.get("arguments"): - tool_calls_acc[idx][ - "function" - ]["arguments"] += func[ + tool_calls_acc[idx]["function"][ "arguments" - ] + ] += func["arguments"] continue # ── Reasoning tokens (bypass buffer) ── - reasoning = delta.get( - "reasoning_content", "" - ) + reasoning = delta.get("reasoning_content", "") if reasoning: reasoning_accum += reasoning if detect_state != _S_DRAINING: @@ -1902,9 +1913,7 @@ class LlamaCppBackend: cleaned = _strip_tool_markup( cumulative_display, ) - if len(cleaned) > len( - _last_emitted - ): + if len(cleaned) > len(_last_emitted): _last_emitted = cleaned yield { "type": "content", @@ -1913,9 +1922,7 @@ class LlamaCppBackend: elif detect_state == _S_BUFFERING: content_buffer += token - stripped_buf = ( - content_buffer.lstrip() - ) + stripped_buf = content_buffer.lstrip() if not stripped_buf: continue @@ -1923,14 +1930,10 @@ class LlamaCppBackend: is_prefix = False is_match = False for sig in _TOOL_XML_SIGNALS: - if stripped_buf.startswith( - sig - ): + if stripped_buf.startswith(sig): is_match = True break - if sig.startswith( - stripped_buf - ): + if sig.startswith(stripped_buf): is_prefix = True break @@ -1946,19 +1949,13 @@ class LlamaCppBackend: # Not a tool -- flush buffer detect_state = _S_STREAMING if in_thinking: - cumulative_display += ( - "" - ) + cumulative_display += "" in_thinking = False - cumulative_display += ( - content_buffer - ) + cumulative_display += content_buffer cleaned = _strip_tool_markup( cumulative_display, ) - if len(cleaned) > len( - _last_emitted - ): + if len(cleaned) > len(_last_emitted): _last_emitted = cleaned yield { "type": "content", @@ -1967,8 +1964,7 @@ class LlamaCppBackend: except json.JSONDecodeError: logger.debug( - f"Skipping malformed SSE line: " - f"{line[:100]}" + f"Skipping malformed SSE line: " f"{line[:100]}" ) if _stream_done: break # exit outer for @@ -1976,8 +1972,10 @@ class LlamaCppBackend: # ── 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 + 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: @@ -1990,7 +1988,8 @@ class LlamaCppBackend: yield { "type": "content", "text": _strip_tool_markup( - cumulative_display, final = True, + cumulative_display, + final = True, ), } else: @@ -2014,27 +2013,19 @@ class LlamaCppBackend: _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 = 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.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) - ) + _mt["predicted_per_second"] = _tn / (_tms / 1000.0) yield { "type": "metadata", "usage": { @@ -2050,6 +2041,7 @@ class LlamaCppBackend: tool_calls = _safety_tc content_text = content_accum import re + content_text = re.sub( r".*?", "", @@ -2083,18 +2075,18 @@ class LlamaCppBackend: 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 = [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".*?", "", @@ -2138,8 +2130,8 @@ class LlamaCppBackend: return # ── Execute tool calls ── - _accumulated_completion_tokens += ( - (_iter_usage or {}).get("completion_tokens", 0) + _accumulated_completion_tokens += (_iter_usage or {}).get( + "completion_tokens", 0 ) _it = _iter_timings or {} _accumulated_predicted_ms += _it.get("predicted_ms", 0) @@ -2170,9 +2162,7 @@ class LlamaCppBackend: 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}" @@ -2235,7 +2225,6 @@ class LlamaCppBackend: return raise - # Clear status yield {"type": "status", "text": ""} @@ -2260,7 +2249,6 @@ class LlamaCppBackend: stream_payload["stop"] = stop stream_payload["stream_options"] = {"include_usage": True} - cumulative = "" _last_emitted = "" in_thinking = False @@ -2272,10 +2260,16 @@ 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 + _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, headers = _auth_headers, + client, + url, + stream_payload, + cancel_event, + headers = _auth_headers, ) as response: if response.status_code != 200: error_body = response.read().decode() @@ -2401,7 +2395,9 @@ class LlamaCppBackend: if not self.is_loaded: return None try: - _auth_headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else {} + _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: @@ -2520,8 +2516,12 @@ class LlamaCppBackend: if need_ids: payload["n_probs"] = 1 - _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: + _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(