TalkToDoc / database.py
goctests0's picture
Upload 2 files
35212b7 verified
Raw
History Blame Contribute Delete
4.68 kB
"""
Database layer for TalkToDoc.
Local development and testing use SQLite (no setup needed). When deployed,
DATABASE_URL is set by Render and the app uses Postgres instead, since
Render's free tier wipes local files like a SQLite database on every
restart.
Queries are written once using SQLite-style '?' placeholders and adapted
automatically for Postgres, so there's a single query per function rather
than two versions of everything.
"""
import os
from pathlib import Path
import sqlite3
DB_PATH = Path(__file__).parent / "talktodoc.db"
SCHEMA_PATH = Path(__file__).parent / "schema.sql"
SCHEMA_PATH_POSTGRES = Path(__file__).parent / "schema_postgres.sql"
DATABASE_URL = os.environ.get("DATABASE_URL")
if DATABASE_URL:
import psycopg2
import psycopg2.extras
def get_connection():
if DATABASE_URL:
return psycopg2.connect(DATABASE_URL, cursor_factory=psycopg2.extras.RealDictCursor)
connection = sqlite3.connect(DB_PATH)
connection.row_factory = sqlite3.Row
connection.execute("PRAGMA foreign_keys = ON")
return connection
def _run(connection, query, params=()):
if DATABASE_URL:
query = query.replace("?", "%s")
cursor = connection.cursor()
cursor.execute(query, params)
return cursor
def _insert_and_get_id(connection, query, params):
if DATABASE_URL:
cursor = _run(connection, query + " RETURNING id", params)
return cursor.fetchone()["id"]
cursor = _run(connection, query, params)
return cursor.lastrowid
def init_db():
schema_path = SCHEMA_PATH_POSTGRES if DATABASE_URL else SCHEMA_PATH
schema_sql = schema_path.read_text()
with get_connection() as connection:
if DATABASE_URL:
connection.cursor().execute(schema_sql)
else:
connection.executescript(schema_sql)
def add_user(name, preferred_language, role):
with get_connection() as connection:
return _insert_and_get_id(
connection,
"INSERT INTO app_user (name, preferred_language, role) VALUES (?, ?, ?)",
(name, preferred_language, role),
)
def add_session(user_id, start_time):
with get_connection() as connection:
return _insert_and_get_id(
connection,
"INSERT INTO session (user_id, start_time) VALUES (?, ?)",
(user_id, start_time),
)
def end_session(session_id, end_time):
with get_connection() as connection:
_run(connection, "UPDATE session SET end_time = ? WHERE id = ?", (end_time, session_id))
def add_interaction(user_id, input_text, detected_language, translated_text, nlu_summary, timestamp):
with get_connection() as connection:
return _insert_and_get_id(
connection,
"""INSERT INTO interaction
(user_id, input_text, detected_language, translated_text, nlu_summary, timestamp)
VALUES (?, ?, ?, ?, ?, ?)""",
(user_id, input_text, detected_language, translated_text, nlu_summary, timestamp),
)
def get_interactions_for_user(user_id):
with get_connection() as connection:
cursor = _run(
connection,
"SELECT * FROM interaction WHERE user_id = ? ORDER BY timestamp",
(user_id,),
)
return [dict(row) for row in cursor.fetchall()]
def get_pending_interactions():
with get_connection() as connection:
cursor = _run(
connection,
"SELECT * FROM interaction WHERE provider_response IS NULL ORDER BY timestamp",
)
return [dict(row) for row in cursor.fetchall()]
def get_completed_interactions():
with get_connection() as connection:
cursor = _run(
connection,
"SELECT * FROM interaction WHERE provider_response IS NOT NULL ORDER BY timestamp DESC",
)
return [dict(row) for row in cursor.fetchall()]
def get_interaction(interaction_id):
with get_connection() as connection:
cursor = _run(connection, "SELECT * FROM interaction WHERE id = ?", (interaction_id,))
row = cursor.fetchone()
return dict(row) if row else None
def update_interaction_response(interaction_id, provider_response, translated_response):
with get_connection() as connection:
_run(
connection,
"UPDATE interaction SET provider_response = ?, translated_response = ? WHERE id = ?",
(provider_response, translated_response, interaction_id),
)
if __name__ == "__main__":
init_db()
if DATABASE_URL:
print("Database initialized (Postgres)")
else:
print("Database created at:", DB_PATH)