"""Connexion PostgreSQL (SQLAlchemy 2)."""
from __future__ import annotations

from collections.abc import Iterator
from contextlib import contextmanager

from sqlalchemy import create_engine
from sqlalchemy.engine import Engine
from sqlalchemy.orm import Session, sessionmaker

from .config import get_config

_engine: Engine | None = None
_factory: sessionmaker[Session] | None = None


def configure(url: str | None = None) -> None:
    """(Ré)initialise le moteur ; les tests l'appellent avec la base de test."""
    global _engine, _factory
    if _engine is not None:
        _engine.dispose()
    _engine = create_engine(
        url or get_config().database_url,
        pool_size=5,
        max_overflow=5,
        pool_pre_ping=True,
        future=True,
    )
    _factory = sessionmaker(_engine, expire_on_commit=False, autoflush=False)


def engine() -> Engine:
    if _engine is None:
        configure()
    assert _engine is not None
    return _engine


def new_session() -> Session:
    if _factory is None:
        configure()
    assert _factory is not None
    return _factory()


def get_db() -> Iterator[Session]:
    db = new_session()
    try:
        yield db
    finally:
        db.close()


@contextmanager
def session_scope() -> Iterator[Session]:
    db = new_session()
    try:
        yield db
        db.commit()
    except Exception:
        db.rollback()
        raise
    finally:
        db.close()
