import os from urllib.parse import quote_plus from sqlalchemy import create_engine, event from sqlalchemy.ext.declarative import declarative_base from sqlalchemy.orm import sessionmaker from dotenv import load_dotenv import datetime from sqlalchemy import inspect as sa_inspect from .utils import audit_log_path, get_current_user_ctx load_dotenv() # Prefer per-component DB config when provided; fallback to DATABASE_URL; finally to sqlite for dev db_driver = os.getenv("DB_DRIVER") db_user = os.getenv("DB_USER") db_password = os.getenv("DB_PASSWORD") db_host = os.getenv("DB_HOST") db_port = os.getenv("DB_PORT") db_name = os.getenv("DB_NAME") database_url = None if db_driver and db_user and db_host and db_port and db_name is not None: # URL-encode username and password. If password is empty or None, omit the ':' portion user_enc = quote_plus(db_user) if db_password: pwd_enc = quote_plus(db_password) auth = f"{user_enc}:{pwd_enc}" else: auth = f"{user_enc}" database_url = f"{db_driver}://{auth}@{db_host}:{db_port}/{db_name}" if not database_url: database_url = os.getenv("DATABASE_URL") if not database_url: # final fallback for local dev database_url = "sqlite:///./test.db" # Engine pool configuration to avoid QueuePool timeouts under concurrency default_pool_size = int(os.getenv("DB_POOL_SIZE", "20")) default_max_overflow = int(os.getenv("DB_POOL_MAX_OVERFLOW", "40")) default_pool_timeout = int(os.getenv("DB_POOL_TIMEOUT", "60")) default_pool_recycle = int(os.getenv("DB_POOL_RECYCLE", "1800")) # seconds if database_url.startswith("sqlite"): engine = create_engine( database_url, pool_pre_ping=True, connect_args={"check_same_thread": False}, ) else: engine = create_engine( database_url, pool_pre_ping=True, pool_size=default_pool_size, max_overflow=default_max_overflow, pool_timeout=default_pool_timeout, pool_recycle=default_pool_recycle, ) SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine) Base = declarative_base() def get_db(): db = SessionLocal() try: yield db finally: db.close() def _write_audit(line: str) -> None: try: with open(audit_log_path(), "a", encoding="utf-8") as f: f.write(line) except Exception: pass @event.listens_for(SessionLocal, "after_flush") def _audit_after_flush(session, context): user = get_current_user_ctx() or {} uid = str((user.get("sub") or user.get("id") or "-")) email = user.get("email") or "-" now = datetime.datetime.utcnow().isoformat() for obj in list(session.new): ins = sa_inspect(obj) mapper = ins.mapper table = mapper.local_table.name if getattr(mapper, "local_table", None) is not None else getattr(obj, "__tablename__", obj.__class__.__name__) pk = ins.identity line = f"{now} | {uid} | {email} | INSERT {table} | pk={pk}\n" _write_audit(line) for obj in list(session.dirty): if not session.is_modified(obj, include_collections=False): continue ins = sa_inspect(obj) mapper = ins.mapper table = mapper.local_table.name if getattr(mapper, "local_table", None) is not None else getattr(obj, "__tablename__", obj.__class__.__name__) changes = [] for attr in ins.attrs: hist = attr.history if hist.has_changes(): old = hist.deleted[0] if hist.deleted else None new = hist.added[0] if hist.added else getattr(obj, attr.key) changes.append(f"{attr.key}={old}->{new}") pk = ins.identity chs = ", ".join(changes) if changes else "-" line = f"{now} | {uid} | {email} | UPDATE {table} | pk={pk} | {chs}\n" _write_audit(line) for obj in list(session.deleted): ins = sa_inspect(obj) mapper = ins.mapper table = mapper.local_table.name if getattr(mapper, "local_table", None) is not None else getattr(obj, "__tablename__", obj.__class__.__name__) pk = ins.identity line = f"{now} | {uid} | {email} | DELETE {table} | pk={pk}\n" _write_audit(line)