File size: 4,419 Bytes
02af5c9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5ae9d9c
 
 
 
 
 
 
 
 
02af5c9
 
 
 
 
 
5ae9d9c
 
 
 
 
 
 
 
 
02af5c9
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
"""Couche base de données — SQLite via SQLAlchemy 2.0.

POC : SQLite, zéro infrastructure. La couche est volontairement isolée
pour pouvoir migrer vers PostgreSQL + PostGIS + pgvector plus tard sans
toucher au reste de l'application.
"""
from collections.abc import Generator

from sqlalchemy import create_engine
from sqlalchemy.orm import DeclarativeBase, Session, sessionmaker

from . import config


class Base(DeclarativeBase):
    pass


# check_same_thread n'est valable que pour SQLite. pool_pre_ping évite les
# connexions Postgres mortes (utile sur Neon/serverless qui se met en veille).
_is_sqlite = config.DATABASE_URL.startswith("sqlite")
engine = create_engine(
    config.DATABASE_URL,
    connect_args={"check_same_thread": False} if _is_sqlite else {},
    pool_pre_ping=not _is_sqlite,
)

# Audit (round 3) : SQLite ignore les FK par défaut → les violations d'intégrité
# référentielle passaient inaperçues en test alors qu'elles sont FATALES sur Neon
# Postgres (NO ACTION). On force PRAGMA foreign_keys=ON pour que les tests
# reproduisent fidèlement la prod et attrapent les régressions de suppression.
if _is_sqlite:
    from sqlalchemy import event

    @event.listens_for(engine, "connect")
    def _sqlite_fk_pragma(dbapi_connection, connection_record):  # noqa: ANN001
        cur = dbapi_connection.cursor()
        cur.execute("PRAGMA foreign_keys=ON")
        cur.close()

SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)


def init_db() -> None:
    """Met la base à jour via Alembic (migrations propres, plus de CHATS_RESET_DB).

    - Base vierge → applique toutes les migrations (crée les tables).
    - Schéma existant sans Alembic (ex. ancienne base create_all) → l'« adopte »
      (stamp) sans tout recréer.
    - Schéma en retard → applique les migrations manquantes (ALTER, etc.).
    CHATS_RESET_DB=1 reste dispo en secours (⚠️ efface tout).
    """
    import os

    from . import models  # noqa: F401  (enregistre les modèles sur Base)

    if os.environ.get("CHATS_RESET_DB", "").lower() in {"1", "true", "yes"}:
        Base.metadata.drop_all(bind=engine)
        _drop_alembic_version()
    _run_migrations()


def _drop_alembic_version() -> None:
    from sqlalchemy import text

    with engine.begin() as conn:
        conn.execute(text("DROP TABLE IF EXISTS alembic_version"))


def _run_migrations() -> None:
    from pathlib import Path

    from alembic import command
    from alembic.config import Config
    from alembic.runtime.migration import MigrationContext
    from sqlalchemy import inspect

    base = Path(__file__).resolve().parent.parent  # .../backend
    cfg = Config(str(base / "alembic.ini"))
    cfg.set_main_option("script_location", str(base / "migrations"))
    try:
        with engine.connect() as conn:
            current = MigrationContext.configure(conn).get_current_revision()
        has_tables = inspect(engine).has_table("cats")
        if current is None and has_tables:
            # Schéma déjà en place mais jamais versionné : c'est forcément une
            # ancienne base create_all figée au schéma INITIAL (787cf3d14385,
            # phone VARCHAR(32)). La stamper directement à « head » lui ferait
            # sauter tous les ALTER intermédiaires (dont l'élargissement de phone
            # en Text). On stampe donc la révision de BASE puis on applique le
            # retard normalement.
            base_rev = _base_revision(cfg) or "787cf3d14385"
            command.stamp(cfg, base_rev)      # adopter le schéma initial
            command.upgrade(cfg, "head")      # rattraper les ALTER manquants
        else:
            command.upgrade(cfg, "head")      # vierge → crée tout ; sinon applique le retard
    except Exception:
        Base.metadata.create_all(bind=engine)  # filet de sécurité


def _base_revision(cfg) -> str | None:  # noqa: ANN001 (Config est un type Alembic)
    """Retourne la première révision de l'historique (down_revision is None)."""
    from alembic.script import ScriptDirectory

    script = ScriptDirectory.from_config(cfg)
    bases = script.get_bases()
    return bases[0] if bases else None


def get_db() -> Generator[Session, None, None]:
    """Dépendance FastAPI : fournit une session et la ferme proprement."""
    db = SessionLocal()
    try:
        yield db
    finally:
        db.close()