Files
audit-web/app/authenticator/database.py
T
2026-08-26 14:11:37 +07:00

115 lines
4.1 KiB
Python

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)