from __future__ import annotations from sqlalchemy import Engine, create_engine from sqlalchemy.pool import StaticPool from cloud.sql_repository import SQLAlchemyCloudRepository def normalize_database_url(database_url: str) -> str: if database_url.startswith("postgresql://"): return "postgresql+psycopg://" + database_url.removeprefix("postgresql://") return database_url def create_database_engine(database_url: str) -> Engine: normalized_url = normalize_database_url(database_url) options: dict[str, object] = {"pool_pre_ping": True} if normalized_url.startswith("sqlite:///"): options["connect_args"] = {"check_same_thread": False} if normalized_url == "sqlite:///:memory:": options["poolclass"] = StaticPool return create_engine(normalized_url, **options) class CloudDatabase: """Own the cloud database engine and repository lifecycle.""" def __init__( self, database_url: str, *, create_schema: bool = True, ) -> None: self.engine = create_database_engine(database_url) self.repository = SQLAlchemyCloudRepository( self.engine, create_schema=create_schema, ) def close(self) -> None: self.repository.close()