SkyAlone / database.py
FreshPixels's picture
Upload 12 files
e683a50 verified
Raw
History Blame
7.67 kB
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:
# Render free tier: 512 MB RAM, keep pool small
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)
# Reverse to chronological order for the LLM
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:
return await conn.fetchval("""
SELECT COUNT(*) FROM messages
WHERE user_id = $1 AND is_summarized = FALSE
""", user_id) 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()