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