Initial commit
This commit is contained in:
@@ -0,0 +1,57 @@
|
||||
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
|
||||
Reference in New Issue
Block a user