115 lines
4.1 KiB
Python
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)
|