Spaces:
Sleeping
Sleeping
| import sqlite3 | |
| import json | |
| import uuid | |
| import time | |
| from datetime import datetime | |
| from pathlib import Path | |
| DB_PATH = Path("/app/data/memory.db") | |
| def _init_db(): | |
| DB_PATH.parent.mkdir(parents=True, exist_ok=True) | |
| conn = sqlite3.connect(str(DB_PATH)) | |
| c = conn.cursor() | |
| c.execute(""" | |
| CREATE TABLE IF NOT EXISTS sessions ( | |
| id TEXT PRIMARY KEY, | |
| title TEXT DEFAULT 'New Chat', | |
| agent TEXT DEFAULT 'researcher', | |
| model TEXT DEFAULT 'gemini-2.0-flash', | |
| created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, | |
| updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP | |
| ) | |
| """) | |
| c.execute(""" | |
| CREATE TABLE IF NOT EXISTS messages ( | |
| id INTEGER PRIMARY KEY AUTOINCREMENT, | |
| session_id TEXT NOT NULL, | |
| role TEXT NOT NULL, | |
| content TEXT NOT NULL, | |
| metadata TEXT DEFAULT '{}', | |
| created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, | |
| FOREIGN KEY (session_id) REFERENCES sessions(id) | |
| ) | |
| """) | |
| conn.commit() | |
| conn.close() | |
| _init_db() | |
| def create_session(agent="researcher", model="gemini-2.0-flash"): | |
| session_id = str(uuid.uuid4())[:8] | |
| conn = sqlite3.connect(str(DB_PATH)) | |
| c = conn.cursor() | |
| c.execute( | |
| "INSERT INTO sessions (id, agent, model, title, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?)", | |
| (session_id, agent, model, "New Chat", datetime.now().isoformat(), datetime.now().isoformat()), | |
| ) | |
| conn.commit() | |
| conn.close() | |
| return session_id | |
| def get_sessions(): | |
| conn = sqlite3.connect(str(DB_PATH)) | |
| conn.row_factory = sqlite3.Row | |
| c = conn.cursor() | |
| c.execute("SELECT * FROM sessions ORDER BY updated_at DESC") | |
| rows = c.fetchall() | |
| conn.close() | |
| return [dict(r) for r in rows] | |
| def get_session(session_id): | |
| conn = sqlite3.connect(str(DB_PATH)) | |
| conn.row_factory = sqlite3.Row | |
| c = conn.cursor() | |
| c.execute("SELECT * FROM sessions WHERE id = ?", (session_id,)) | |
| row = c.fetchone() | |
| conn.close() | |
| return dict(row) if row else None | |
| def delete_session(session_id): | |
| conn = sqlite3.connect(str(DB_PATH)) | |
| c = conn.cursor() | |
| c.execute("DELETE FROM messages WHERE session_id = ?", (session_id,)) | |
| c.execute("DELETE FROM sessions WHERE id = ?", (session_id,)) | |
| conn.commit() | |
| conn.close() | |
| def add_message(session_id, role, content, metadata=None): | |
| conn = sqlite3.connect(str(DB_PATH)) | |
| c = conn.cursor() | |
| meta_json = json.dumps(metadata or {}, ensure_ascii=False) | |
| c.execute( | |
| "INSERT INTO messages (session_id, role, content, metadata) VALUES (?, ?, ?, ?)", | |
| (session_id, role, content, meta_json), | |
| ) | |
| # Update session timestamp | |
| c.execute( | |
| "UPDATE sessions SET updated_at = ? WHERE id = ?", | |
| (datetime.now().isoformat(), session_id), | |
| ) | |
| conn.commit() | |
| conn.close() | |
| def get_messages(session_id, limit=20): | |
| conn = sqlite3.connect(str(DB_PATH)) | |
| conn.row_factory = sqlite3.Row | |
| c = conn.cursor() | |
| c.execute( | |
| "SELECT * FROM messages WHERE session_id = ? ORDER BY created_at DESC LIMIT ?", | |
| (session_id, limit), | |
| ) | |
| rows = c.fetchall() | |
| conn.close() | |
| messages = [dict(r) for r in reversed(rows)] | |
| for m in messages: | |
| if isinstance(m.get("metadata"), str): | |
| try: | |
| m["metadata"] = json.loads(m["metadata"]) | |
| except json.JSONDecodeError: | |
| pass | |
| return messages | |
| def update_session_title(session_id, title): | |
| conn = sqlite3.connect(str(DB_PATH)) | |
| c = conn.cursor() | |
| c.execute("UPDATE sessions SET title = ? WHERE id = ?", (title[:60], session_id)) | |
| conn.commit() | |
| conn.close() | |