"""Database session and methods for lambda integration tests.""" import logging from typing import Any from typing import cast from sqlalchemy import column from sqlalchemy import create_engine from sqlalchemy import delete from sqlalchemy import Engine from sqlalchemy import insert from sqlalchemy import literal_column from sqlalchemy import Result from sqlalchemy import Select from sqlalchemy import select from sqlalchemy import table from sqlalchemy import text from sqlalchemy import update from sqlalchemy.engine import CursorResult from sqlalchemy.engine import URL from sqlalchemy.orm import scoped_session from sqlalchemy.orm import Session from sqlalchemy.orm import sessionmaker from tests import config from tests.src.secrets_manager import get_secret logger = logging.getLogger(__name__) def _build_engine() -> Engine: """Create and return a SQLAlchemy engine using credentials from Secrets Manager.""" db_name = config.DB_NAME credentials: dict[str, Any] = get_secret(config.DB_SECRET_NAME) url = URL.create( drivername='mysql+pymysql', username=credentials['user'], password=credentials['password'], host=credentials['host'], port=credentials['port'], database=db_name, ) return create_engine( url, pool_pre_ping=True, pool_recycle=3600, echo=False, # Set to True only for debugging # MySQL defaults to REPEATABLE READ which causes stale reads. # READ COMMITTED ensures we see committed changes from other transactions. isolation_level='READ COMMITTED', ) _engine: Engine | None = None _session_factory: scoped_session[Session] | None = None def get_session_factory() -> scoped_session[Session]: """Return the module-level scoped session factory, initializing it on first call.""" global _engine, _session_factory if _session_factory is None: _engine = _build_engine() _session_factory = scoped_session(sessionmaker(bind=_engine)) return _session_factory def execute_query( db_session: Session, query: str, params: dict[str, Any] ) -> Result[Any]: """Execute a raw SQL query and return the SQLAlchemy Result object.""" return db_session.execute(text(query), params) def _map_one(result: Result[Any]) -> dict[str, Any] | None: """Return the first row as a dict, or None when no rows are found.""" row = result.mappings().fetchone() if row: return dict(row) logger.debug('No row matched the given conditions.') return None def get_entity( db_session: Session, table_name: str, conditions: dict[str, Any], ) -> dict[str, Any] | None: """Fetch a single row (all columns) from *table_name* matching *conditions*.""" table_obj = table(table_name, *(column(column_name) for column_name in conditions)) stmt: Select[tuple[Any]] = ( select(literal_column('*')) .select_from(table_obj) .where( *( table_obj.c[column_name] == value for column_name, value in conditions.items() ) ) ) return _map_one(db_session.execute(stmt)) def update_entity( db_session: Session, table_name: str, conditions: dict[str, Any], values: dict[str, Any], ) -> None: """Update rows in *table_name* matching *conditions* with *values*.""" table_obj = table( table_name, *(column(column_name) for column_name in {**conditions, **values}) ) stmt = ( update(table_obj) .where( *( table_obj.c[column_name] == value for column_name, value in conditions.items() ) ) .values(**values) ) db_session.execute(stmt) def insert_entity( db_session: Session, table_name: str, values: dict[str, Any], id_column: str, delete_conditions: dict[str, Any] | None = None, ) -> dict[str, Any] | None: """Insert a single row into *table_name* and return the full inserted row as a dict. If *delete_conditions* are provided, any existing rows matching them are deleted first. """ if delete_conditions: delete_table = table( table_name, *(column(column_name) for column_name in delete_conditions) ) db_session.execute( delete(delete_table).where( *( delete_table.c[column_name] == value for column_name, value in delete_conditions.items() ) ) ) table_obj = table(table_name, *(column(column_name) for column_name in values)) cursor_result = cast( CursorResult[Any], db_session.execute(insert(table_obj).values(**values)) ) table_obj = table(table_name, column(id_column)) inserted_row_stmt: Select[Any] = ( select(literal_column('*')) .select_from(table_obj) .where(table_obj.c[id_column] == cursor_result.lastrowid) ) return _map_one(db_session.execute(inserted_row_stmt)) def delete_entity_by_id( db_session: Session, table_name: str, id_column: str, entity_id: int, ) -> None: """Delete row(s) from *table_name* by *entity_id*.""" table_obj = table(table_name, column(id_column)) db_session.execute(delete(table_obj).where(table_obj.c[id_column] == entity_id))