Spaces:
Runtime error
Runtime error
File size: 7,575 Bytes
e683a50 d1983a4 e683a50 d1983a4 e683a50 d1983a4 e683a50 | 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 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 | import asyncpg
import logging
from typing import Optional, List, Dict, Any
from config import config
logger = logging.getLogger(__name__)
class Database:
def __init__(self) -> None:
self.pool: Optional[asyncpg.Pool] = None
async def connect(self) -> None:
self.pool = await asyncpg.create_pool(
dsn=config.DATABASE_URL,
min_size=1,
max_size=3,
command_timeout=60,
)
logger.info("Database pool created (max_size=3)")
await self._create_tables()
async def disconnect(self) -> None:
if self.pool:
await self.pool.close()
logger.info("Database pool closed")
def _acquire(self):
if self.pool is None:
raise RuntimeError("Database not connected. Call connect() first.")
return self.pool.acquire()
async def _create_tables(self) -> None:
async with self._acquire() as conn:
await conn.execute("""
CREATE TABLE IF NOT EXISTS users (
id BIGINT PRIMARY KEY,
username VARCHAR(255),
first_name VARCHAR(255),
last_name VARCHAR(255),
created_at TIMESTAMP WITH TIME ZONE DEFAULT NOW(),
updated_at TIMESTAMP WITH TIME ZONE DEFAULT NOW()
)
""")
await conn.execute("""
CREATE TABLE IF NOT EXISTS messages (
id SERIAL PRIMARY KEY,
user_id BIGINT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
role VARCHAR(20) NOT NULL CHECK (role IN ('user', 'assistant', 'system')),
content TEXT NOT NULL,
is_summarized BOOLEAN DEFAULT FALSE,
created_at TIMESTAMP WITH TIME ZONE DEFAULT NOW()
)
""")
await conn.execute("""
CREATE INDEX IF NOT EXISTS idx_messages_user_id_created_at
ON messages(user_id, created_at DESC)
""")
await conn.execute("""
CREATE INDEX IF NOT EXISTS idx_messages_user_id_summarized
ON messages(user_id, is_summarized, created_at DESC)
""")
await conn.execute("""
CREATE TABLE IF NOT EXISTS summaries (
id SERIAL PRIMARY KEY,
user_id BIGINT NOT NULL UNIQUE REFERENCES users(id) ON DELETE CASCADE,
summary TEXT NOT NULL,
message_count INTEGER NOT NULL DEFAULT 0,
created_at TIMESTAMP WITH TIME ZONE DEFAULT NOW(),
updated_at TIMESTAMP WITH TIME ZONE DEFAULT NOW()
)
""")
logger.info("Database tables created/verified")
async def upsert_user(
self,
user_id: int,
username: Optional[str],
first_name: Optional[str],
last_name: Optional[str],
) -> None:
async with self._acquire() as conn:
await conn.execute("""
INSERT INTO users (id, username, first_name, last_name)
VALUES ($1, $2, $3, $4)
ON CONFLICT (id) DO UPDATE SET
username = EXCLUDED.username,
first_name = EXCLUDED.first_name,
last_name = EXCLUDED.last_name,
updated_at = NOW()
""", user_id, username, first_name, last_name)
async def save_message(self, user_id: int, role: str, content: str) -> None:
async with self._acquire() as conn:
await conn.execute("""
INSERT INTO messages (user_id, role, content)
VALUES ($1, $2, $3)
""", user_id, role, content)
async def get_messages(self, user_id: int, limit: int = 30) -> List[Dict[str, Any]]:
async with self._acquire() as conn:
rows = await conn.fetch("""
SELECT role, content, created_at
FROM messages
WHERE user_id = $1 AND is_summarized = FALSE
ORDER BY created_at DESC
LIMIT $2
""", user_id, limit)
return [
{"role": r["role"], "content": r["content"], "created_at": r["created_at"]}
for r in reversed(rows)
]
async def get_summary(self, user_id: int) -> Optional[str]:
async with self._acquire() as conn:
row = await conn.fetchrow("""
SELECT summary FROM summaries WHERE user_id = $1
""", user_id)
return row["summary"] if row else None
async def save_summary(self, user_id: int, summary: str, message_count: int) -> None:
async with self._acquire() as conn:
await conn.execute("""
INSERT INTO summaries (user_id, summary, message_count, updated_at)
VALUES ($1, $2, $3, NOW())
ON CONFLICT (user_id) DO UPDATE SET
summary = EXCLUDED.summary,
message_count = summaries.message_count + EXCLUDED.message_count,
updated_at = NOW()
""", user_id, summary, message_count)
async def mark_summarized(self, user_id: int, cutoff_id: int) -> None:
async with self._acquire() as conn:
await conn.execute("""
UPDATE messages
SET is_summarized = TRUE
WHERE user_id = $1 AND id <= $2
""", user_id, cutoff_id)
async def get_oldest_unsummarized(self, user_id: int, limit: int) -> List[Dict[str, Any]]:
async with self._acquire() as conn:
rows = await conn.fetch("""
SELECT id, role, content
FROM messages
WHERE user_id = $1 AND is_summarized = FALSE
ORDER BY created_at ASC
LIMIT $2
""", user_id, limit)
return [{"id": r["id"], "role": r["role"], "content": r["content"]} for r in rows]
async def count_unsummarized(self, user_id: int) -> int:
async with self._acquire() as conn:
val = await conn.fetchval("""
SELECT COUNT(*) FROM messages
WHERE user_id = $1 AND is_summarized = FALSE
""", user_id)
return val or 0
async def clear_history(self, user_id: int) -> int:
async with self._acquire() as conn:
result = await conn.execute("""
DELETE FROM messages WHERE user_id = $1
""", user_id)
await conn.execute("""
DELETE FROM summaries WHERE user_id = $1
""", user_id)
try:
count = int(result.split()[-1])
except (ValueError, IndexError):
count = 0
logger.info("Cleared %d messages and summary for user %s", count, user_id)
return count
async def get_stats(self, user_id: int) -> Dict[str, Any]:
async with self._acquire() as conn:
user_count = await conn.fetchval("SELECT COUNT(*) FROM users")
msg_count = await conn.fetchval(
"SELECT COUNT(*) FROM messages WHERE user_id = $1", user_id
)
total_msg_count = await conn.fetchval("SELECT COUNT(*) FROM messages")
return {
"total_users": user_count,
"user_messages": msg_count,
"total_messages": total_msg_count,
}
db = Database()
|