import os import threading import httpx MODAL_ENDPOINT = os.environ.get("MODAL_ENDPOINT_LLM", "") MODAL_AUTH_TOKEN = os.environ.get("MODAL_AUTH_TOKEN", "") TRUNCATION_LIMIT = 20_000 def warmup(): """Fire-and-forget request to pre-load the LLM into GPU memory. Called once on page load from reset_session(). The Modal container auto-scales down after scaledown_window (90s) of inactivity. """ if not MODAL_ENDPOINT: return def _ping(): try: httpx.post( MODAL_ENDPOINT + "/warmup", json={"token": MODAL_AUTH_TOKEN}, timeout=120, ) except Exception: pass threading.Thread(target=_ping, daemon=True).start() def chat(messages: list[dict]) -> str | None: if not MODAL_ENDPOINT: raise RuntimeError("MODAL_ENDPOINT_LLM not set") _trim(messages) resp = httpx.post( MODAL_ENDPOINT, json={"messages": messages, "token": MODAL_AUTH_TOKEN}, timeout=600, ) resp.raise_for_status() return resp.json().get("text") def _trim(messages: list[dict]): if len(messages) < 4: return total_chars = sum(len(m.get("content", "")) for m in messages) if total_chars < TRUNCATION_LIMIT * 3.5: return system = [m for m in messages if m.get("role") == "system"] rest = [m for m in messages if m.get("role") != "system"] while rest and total_chars >= TRUNCATION_LIMIT * 3.5: dropped = rest.pop(0) total_chars -= len(dropped.get("content", "")) messages[:] = system + rest