study-buddy / app /agents /cerebras_client.py
GitHub Actions
deploy d092bea3608b7a29952f16357fda39b7a29e399b
2e818da
Raw
History Blame Contribute Delete
19.2 kB
"""Cerebras SDK wrapper.
Responsibilities:
- structured_complete(): call API with strict JSON schema, parse into a Pydantic model
- stream_complete(): async generator yielding text tokens
- Rate-limit short-circuit: tracks cooldown window, raises CerebrasError immediately
- Health mailbox: get_health() exposes current state for the /api/health endpoint
- Schema builder: strips $defs and enforces additionalProperties=false at all levels
Each method is self-contained. If structured_complete breaks, stream still works.
"""
import json
import os
import time
import threading
from typing import Any, AsyncIterator, Dict, List, Optional, Type
from pydantic import BaseModel, ValidationError
from app.agents.cerebras_errors import CerebrasError, CerebrasErrorKind, classify_error
from app.observability.operation import observe_operation
MODEL_ID = "gemma-4-31b"
_deployment_env = os.getenv("DEPLOYMENT_ENV", "desktop")
_cerebras_concurrency = 5 if _deployment_env == "demo" else 50
_cerebras_semaphore = threading.Semaphore(_cerebras_concurrency)
class CerebrasClient:
def __init__(self, api_key: Optional[str] = None) -> None:
from cerebras.cloud.sdk import Cerebras
self._client = Cerebras(api_key=api_key or os.getenv("CEREBRAS_API_KEY"))
self._health: dict = {"status": "ok"}
self._rate_limit_until: float = 0.0
# ------------------------------------------------------------------ #
# Health #
# ------------------------------------------------------------------ #
def get_health(self) -> dict:
return dict(self._health)
# ------------------------------------------------------------------ #
# Internal helpers #
# ------------------------------------------------------------------ #
def _raise_if_cooling_down(self) -> None:
"""Short-circuit a call made while a rate-limit cooldown window is still
active -> raise RATE_LIMITED immediately instead of blocking a worker for
up to the cooldown and then hitting the API (which would just 429 again).
Matches the documented contract: 'tracks rate-limit cooldown windows and
short-circuits future calls without hitting the API.'
"""
remaining = self._rate_limit_until - time.time()
if remaining > 0:
err = CerebrasError(
kind=CerebrasErrorKind.RATE_LIMITED,
message=f"Rate limit cooldown active -> retry in {remaining:.0f}s.",
retry_after_seconds=int(remaining) + 1,
)
self._health = {"status": "error", **err.to_frontend_payload()}
raise err
async def _async_check_rate_limit(self) -> None:
self._raise_if_cooling_down()
def _sync_check_rate_limit(self) -> None:
self._raise_if_cooling_down()
def _handle_sdk_exc(self, exc: Exception, op: Any = None) -> None:
err = classify_error(exc)
if err.kind == CerebrasErrorKind.RATE_LIMITED and err.retry_after_seconds:
self._rate_limit_until = time.time() + err.retry_after_seconds
self._health = {"status": "error", **err.to_frontend_payload()}
if op is not None:
op.mark_error(err.kind.value)
raise err from exc
def _log_usage(self, project_id: str, operation: str, tokens: int | float, elapsed: float, model: str) -> None:
if not project_id or not isinstance(tokens, (int, float)) or not tokens or elapsed <= 0:
return
try:
from app.services.project_activity import ProjectActivityService
ProjectActivityService().record_llm_usage(project_id, operation, tokens, elapsed, model)
except Exception:
pass
def _build_schema(self, model: Type[BaseModel]) -> Dict[str, Any]:
raw = model.model_json_schema()
defs = raw.pop("$defs", {})
def _inline(obj: Any, path: frozenset = frozenset()) -> Any:
"""Recursively replace every {"$ref": "#/$defs/Name"} with a copy of
the referenced definition. Cerebras strict mode rejects $ref, so the
schema must be fully self-contained.
`path` tracks which $defs names are currently being expanded along
this branch, to detect a self-referential/recursive Pydantic model
(e.g. a tree node with children of its own type) -- inlining one of
those with no cycle check recurses forever (confirmed: RecursionError,
hit by HierNode before it was redesigned -- see
app/agents/d3/templates/_hierarchy.py). Cerebras's strict mode can't
represent true recursion anyway (it rejects $ref outright), so a
genuinely cyclic schema needs a bounded-depth redesign, not a runtime
workaround -- this raises a clear, immediate error naming the cycle
instead of silently blowing the stack.
"""
if isinstance(obj, list):
return [_inline(v, path) for v in obj]
if not isinstance(obj, dict):
return obj
ref = obj.get("$ref")
if isinstance(ref, str) and ref.startswith("#/$defs/"):
name = ref.split("/")[-1]
if name in path:
raise ValueError(
f"Cerebras schema build hit a self-referential/recursive model "
f"({name!r}) -- strict json_schema mode cannot represent unbounded "
f"recursion. Redesign it as a fixed-depth chain of distinct classes "
f"instead (see HierNode/HierNodeL2/HierNodeL3/HierLeaf in "
f"app/agents/d3/templates/_hierarchy.py for the pattern)."
)
target = defs.get(name, {})
merged = {**target, **{k: v for k, v in obj.items() if k != "$ref"}}
return _inline(merged, path | {name})
return {k: _inline(v, path) for k, v in obj.items()}
raw = _inline(raw)
raw["additionalProperties"] = False
def _fix(obj: Any) -> None:
if not isinstance(obj, dict):
return
if obj.get("type") == "object":
if "properties" not in obj and "anyOf" not in obj and isinstance(obj.get("additionalProperties"), dict):
# An open-ended dict[str, X] field -> Cerebras strict json_schema mode rejects
# any object schema that isn't fully enumerated ("Object fields require at least
# one of: 'properties' or 'anyOf'"). This only ever surfaces as a confusing 400
# from the API well after the request is built, so fail immediately and name the
# field instead. Reshape dict[str, X] into a list of an explicit {key, value} model.
raise ValueError(
f"Cerebras schema build hit an open-ended dict field ({obj.get('title', 'unnamed field')!r}) "
"-- strict json_schema mode requires every object schema to declare 'properties' or "
"'anyOf'. Replace the dict[str, X] field with a list of an explicit Pydantic model "
"(e.g. list[KeyValue] with key/value fields) instead."
)
obj.setdefault("additionalProperties", False)
if obj.get("type") == "string":
# Cerebras strict json_schema mode rejects minLength/maxLength on
# string types ("Invalid fields for schema with types ['string']").
obj.pop("minLength", None)
obj.pop("maxLength", None)
if obj.get("type") == "array":
# Cerebras likewise rejects JSON Schema cardinality keywords on
# arrays. Pydantic still enforces these constraints after the
# response is parsed; only the provider-facing schema is relaxed.
obj.pop("minItems", None)
obj.pop("maxItems", None)
for v in obj.values():
if isinstance(v, list):
for item in v:
_fix(item)
else:
_fix(v)
_fix(raw)
return raw
# ------------------------------------------------------------------ #
# Public API #
# ------------------------------------------------------------------ #
def structured_complete(
self, messages: List[Dict[str, Any]], output_model: Type[BaseModel], model: str | None = None, **kwargs
) -> BaseModel:
from cerebras.cloud.sdk import APIConnectionError as _ConnErr
project_id = str(kwargs.pop("project_id", "") or "")
trace_label = str(kwargs.pop("trace_label", output_model.__name__) or output_model.__name__)
with observe_operation(
"llm.generate",
subsystem="llm",
consumer="structured",
attributes={"llm_model": model or MODEL_ID, "llm_schema": output_model.__name__},
) as op:
self._sync_check_rate_limit()
schema = self._build_schema(output_model)
response_format = {
"type": "json_schema",
"json_schema": {
"name": output_model.__name__,
"strict": True,
"schema": schema,
},
}
last_exc: Exception = RuntimeError("unreachable")
max_retries = 3
print(f"[Cerebras] structured_complete queued trace={trace_label} schema={output_model.__name__}", flush=True)
for attempt in range(max_retries):
self._sync_check_rate_limit()
try:
t0 = time.time()
with _cerebras_semaphore:
print(
f"[Cerebras] structured_complete dispatch trace={trace_label} attempt={attempt + 1}/{max_retries}",
flush=True,
)
resp = self._client.chat.completions.create(
model=model or MODEL_ID, messages=messages, response_format=response_format, **kwargs
)
t1 = time.time()
tokens = getattr(resp.usage, "completion_tokens", 0)
if isinstance(tokens, (int, float)) and tokens and (t1 - t0) > 0:
suffix = " (retry)" if attempt > 0 else ""
print(f"[Cerebras Speed] structured_complete{suffix} trace={trace_label} generated {tokens} tokens in {t1-t0:.2f}s ({tokens/(t1-t0):.1f} tok/s)", flush=True)
self._log_usage(project_id, "structured_complete", tokens, t1 - t0, model or MODEL_ID)
op.add_count("llm_tokens", int(tokens))
raw = resp.choices[0].message.content
self._health = {"status": "ok"}
op.set("llm_attempts", attempt + 1)
return output_model.model_validate_json(raw)
except _ConnErr as exc:
# HTTP/2 server disconnect -> retry once with backoff
print(f"[Cerebras] connection error trace={trace_label} attempt={attempt + 1}; retrying", flush=True)
last_exc = exc
time.sleep(2)
continue
except Exception as exc:
err = classify_error(exc)
if err.kind == CerebrasErrorKind.RATE_LIMITED and attempt < max_retries - 1:
delay = err.retry_after_seconds or (2 ** attempt)
print(f"[Cerebras] Rate limited. Retrying structured_complete in {delay}s...")
time.sleep(delay)
last_exc = exc
continue
if isinstance(exc, (json.JSONDecodeError, ValidationError)):
print(f"[Cerebras] structured response validation failed trace={trace_label} attempt={attempt + 1}", flush=True)
last_exc = exc
if attempt == 0:
continue
self._handle_sdk_exc(exc, op)
self._handle_sdk_exc(last_exc, op)
def complete_with_tools(
self,
messages: List[Dict[str, Any]],
tools: List[Dict[str, Any]],
tool_choice: str = "auto",
model: str | None = None,
**kwargs
):
"""Non-streaming completion that may return tool calls.
Returns the raw SDK message object (has `.content` and `.tool_calls`).
The caller inspects `.tool_calls`, executes them, appends the results,
then streams the final answer via stream_complete().
"""
max_retries = 3
last_exc: Exception = RuntimeError("unreachable")
project_id = str(kwargs.pop("project_id", "") or "")
with observe_operation(
"llm.generate", subsystem="llm", consumer="tools",
attributes={"llm_model": model or MODEL_ID},
) as op:
for attempt in range(max_retries):
self._sync_check_rate_limit()
try:
t0 = time.time()
with _cerebras_semaphore:
resp = self._client.chat.completions.create(
model=model or MODEL_ID,
messages=messages,
tools=tools,
tool_choice=tool_choice,
**kwargs
)
t1 = time.time()
tokens = getattr(resp.usage, "completion_tokens", 0)
if isinstance(tokens, (int, float)) and tokens and (t1 - t0) > 0:
print(f"[Cerebras Speed] complete_with_tools generated {tokens} tokens in {t1-t0:.2f}s ({tokens/(t1-t0):.1f} tok/s)")
self._log_usage(project_id, "complete_with_tools", tokens, t1 - t0, model or MODEL_ID)
op.add_count("llm_tokens", int(tokens))
self._health = {"status": "ok"}
op.set("llm_attempts", attempt + 1)
return resp.choices[0].message
except Exception as exc:
err = classify_error(exc)
if err.kind == CerebrasErrorKind.RATE_LIMITED and attempt < max_retries - 1:
delay = err.retry_after_seconds or (2 ** attempt)
print(f"[Cerebras] Rate limited. Retrying complete_with_tools in {delay}s...")
time.sleep(delay)
last_exc = exc
continue
self._handle_sdk_exc(exc, op)
self._handle_sdk_exc(last_exc, op)
async def stream_complete(self, messages: List[Dict[str, Any]], model: str | None = None, **kwargs) -> AsyncIterator[str]:
max_retries = 3
last_exc: Exception = RuntimeError("unreachable")
project_id = str(kwargs.pop("project_id", "") or "")
with observe_operation(
"llm.generate", subsystem="llm", consumer="stream",
attributes={"llm_model": model or MODEL_ID, "llm_streaming": True},
) as op:
for attempt in range(max_retries):
await self._async_check_rate_limit()
try:
t0 = time.time()
first_token_at: float | None = None
with _cerebras_semaphore:
# Note: this blocks the thread during generator creation.
# We only hold the semaphore for the initial call, the streaming
# generator yields responses over time. But wait, if stream is sync,
# the generator yields in real-time. Since we process chunks outside the
# semaphore, we don't hold the semaphore while yielding.
stream = self._client.chat.completions.create(
model=model or MODEL_ID, messages=messages, stream=True, **kwargs
)
chunk_count = 0
usage_logged = False
for chunk in stream:
usage = getattr(chunk, "usage", None)
delta = chunk.choices[0].delta.content if chunk.choices else None
if delta:
if first_token_at is None:
first_token_at = time.time()
ttft_ms = round((first_token_at - t0) * 1000.0, 3)
op.set("llm_ttft_ms", ttft_ms)
op.event("llm_first_token", {"llm_ttft_ms": ttft_ms})
chunk_count += 1
yield delta
if usage and getattr(usage, "completion_tokens", 0) and not usage_logged:
t1 = time.time()
tokens = usage.completion_tokens
if isinstance(tokens, (int, float)) and tokens and (t1 - t0) > 0:
print(f"[Cerebras Speed] stream_complete generated {tokens} tokens in {t1-t0:.2f}s ({tokens/(t1-t0):.1f} tok/s)")
self._log_usage(project_id, "stream_complete", tokens, t1 - t0, model or MODEL_ID)
op.add_count("llm_tokens", int(tokens))
usage_logged = True
if not usage_logged:
t1 = time.time()
if (t1 - t0) > 0 and chunk_count > 0:
print(f"[Cerebras Speed] stream_complete (approx) generated {chunk_count} chunks in {t1-t0:.2f}s ({chunk_count/(t1-t0):.1f} chunks/s)")
self._log_usage(project_id, "stream_complete_approx", chunk_count, t1 - t0, model or MODEL_ID)
op.add_count("llm_chunks_approx", chunk_count)
self._health = {"status": "ok"}
op.set("llm_attempts", attempt + 1)
return
except Exception as exc:
err = classify_error(exc)
if err.kind == CerebrasErrorKind.RATE_LIMITED and attempt < max_retries - 1:
delay = err.retry_after_seconds or (2 ** attempt)
print(f"[Cerebras] Rate limited. Retrying stream_complete in {delay}s...")
import asyncio
await asyncio.sleep(delay)
last_exc = exc
continue
self._handle_sdk_exc(exc, op)
self._handle_sdk_exc(last_exc, op)