58 lines
1.9 KiB
Python
58 lines
1.9 KiB
Python
from enum import Enum
|
|
|
|
from sqlalchemy import inspect, text
|
|
from sqlmodel import Session, SQLModel, create_engine
|
|
|
|
from app.core.config import get_settings
|
|
|
|
settings = get_settings()
|
|
|
|
connect_args = {"check_same_thread": False} if settings.database_url.startswith("sqlite") else {}
|
|
engine = create_engine(settings.database_url, connect_args=connect_args)
|
|
|
|
|
|
def _default_sql(column) -> str:
|
|
if column.default is None or not getattr(column.default, "is_scalar", False):
|
|
return ""
|
|
value = column.default.arg
|
|
if callable(value):
|
|
return ""
|
|
if isinstance(value, Enum):
|
|
value = value.value
|
|
if isinstance(value, str):
|
|
return f" DEFAULT '{value.replace(chr(39), chr(39) * 2)}'"
|
|
if isinstance(value, bool):
|
|
return f" DEFAULT {int(value)}"
|
|
return f" DEFAULT {value}"
|
|
|
|
|
|
def _add_missing_columns() -> None:
|
|
"""Additive, best-effort schema sync for SQLite: adds columns that exist on the
|
|
SQLModel definitions but not yet in the DB, so new fields don't require wiping
|
|
the database. Does not handle renames, type changes, or column removal."""
|
|
inspector = inspect(engine)
|
|
existing_tables = set(inspector.get_table_names())
|
|
|
|
with engine.begin() as conn:
|
|
for table in SQLModel.metadata.sorted_tables:
|
|
if table.name not in existing_tables:
|
|
continue
|
|
existing_columns = {col["name"] for col in inspector.get_columns(table.name)}
|
|
for column in table.columns:
|
|
if column.name in existing_columns:
|
|
continue
|
|
col_type = column.type.compile(dialect=engine.dialect)
|
|
conn.execute(
|
|
text(f'ALTER TABLE "{table.name}" ADD COLUMN "{column.name}" {col_type}{_default_sql(column)}')
|
|
)
|
|
|
|
|
|
def create_db_and_tables() -> None:
|
|
SQLModel.metadata.create_all(engine)
|
|
_add_missing_columns()
|
|
|
|
|
|
def get_session():
|
|
with Session(engine) as session:
|
|
yield session
|