File size: 1,574 Bytes
88366bb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import os
import asyncpg

_pool: asyncpg.Pool | None = None

DDL = """
CREATE TABLE IF NOT EXISTS messages (
    id BIGSERIAL PRIMARY KEY,
    session_id TEXT NOT NULL,
    role TEXT NOT NULL,          -- 'user' | 'assistant' | 'tool'
    content TEXT NOT NULL,
    created_at TIMESTAMPTZ NOT NULL DEFAULT now()
);
CREATE INDEX IF NOT EXISTS idx_messages_session ON messages (session_id, created_at);
"""


async def get_pool() -> asyncpg.Pool:
    global _pool
    if _pool is None:
        dsn = os.environ["NEON_DATABASE_URL"]  # e.g. postgresql://user:pass@ep-xxx.neon.tech/db?sslmode=require
        # asyncpg wants sslmode stripped, ssl handled separately
        dsn = dsn.replace("?sslmode=require", "").replace("&sslmode=require", "")
        _pool = await asyncpg.create_pool(dsn=dsn, ssl="require", min_size=1, max_size=5)
        async with _pool.acquire() as conn:
            await conn.execute(DDL)
    return _pool


async def load_history(session_id: str, limit: int = 20) -> list[dict]:
    pool = await get_pool()
    rows = await pool.fetch(
        """
        SELECT role, content FROM messages
        WHERE session_id = $1
        ORDER BY created_at DESC
        LIMIT $2
        """,
        session_id, limit,
    )
    return [{"role": r["role"], "content": r["content"]} for r in reversed(rows)]


async def save_message(session_id: str, role: str, content: str) -> None:
    pool = await get_pool()
    await pool.execute(
        "INSERT INTO messages (session_id, role, content) VALUES ($1, $2, $3)",
        session_id, role, content,
    )