File size: 1,607 Bytes
b05542b eb93c12 b05542b f234692 b05542b 4ff45b3 b05542b eb93c12 b18c1b0 eb93c12 b18c1b0 eb93c12 b18c1b0 eb93c12 b05542b f234692 b05542b 4ff45b3 b05542b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 | 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
|