diff --git a/src/transcription/config.py b/src/transcription/config.py index 75c857f..f3aa56f 100644 --- a/src/transcription/config.py +++ b/src/transcription/config.py @@ -41,6 +41,7 @@ class Settings(BaseSettings): # --- persistence --- database_url: str = "sqlite:///./transcription.db" bootstrap_schema_on_startup: bool | None = None + sqlite_check_same_thread: bool = False # --- filesystem paths --- upload_dir: Path = Path("./uploads") @@ -61,10 +62,10 @@ class Settings(BaseSettings): _settings: ContextVar[Settings | None] = ContextVar("settings", default=None) -def get_settings() -> Settings: +def get_settings(**kwargs) -> Settings: settings = _settings.get() if settings is None: - settings = Settings() # pyright: ignore[reportCallIssue] + settings = Settings(**kwargs) # pyright: ignore[reportCallIssue] _settings.set(settings) return settings diff --git a/src/transcription/db/runtime.py b/src/transcription/db/runtime.py index a0905f4..cf6f6c0 100644 --- a/src/transcription/db/runtime.py +++ b/src/transcription/db/runtime.py @@ -79,24 +79,25 @@ def initialize_database_runtime(*, settings: Settings | None = None) -> Database return runtime -def get_engine() -> AsyncEngine: +def get_engine(settings: Settings | None = None) -> AsyncEngine: """Return the current async SQLAlchemy engine.""" - runtime = _runtime.get() or initialize_database_runtime() + runtime = _runtime.get() or initialize_database_runtime(settings=settings) return runtime.engine -def get_session_factory() -> async_sessionmaker[AsyncSession]: +def get_session_factory(settings: Settings | None = None) -> async_sessionmaker[AsyncSession]: """Return the shared async session factory.""" - runtime = _runtime.get() or initialize_database_runtime() + runtime = _runtime.get() or initialize_database_runtime(settings=settings) return runtime.session_factory @asynccontextmanager async def get_session( *, + settings: Settings | None = None, session_factory: async_sessionmaker[AsyncSession] | None = None, ) -> AsyncGenerator[AsyncSession]: """Yield a database session and ensure cleanup.""" - active_session_factory = session_factory or get_session_factory() + active_session_factory = session_factory or get_session_factory(settings) async with active_session_factory() as session: yield session diff --git a/tests/conftest.py b/tests/conftest.py index 1d6f4b0..da32855 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -11,7 +11,11 @@ from sqlmodel import SQLModel from sqlmodel import create_engine from sqlmodel.pool import StaticPool +from transcription.config import Settings +from transcription.config import get_settings +from transcription.db.operations import create_all from transcription.db.runtime import dispose_database_runtime +from transcription.db.runtime import get_engine from transcription.db.runtime import get_session @@ -29,9 +33,17 @@ def session(): @pytest_asyncio.fixture -async def async_session(): +async def default_settings(): + """Provide default settings for tests.""" + settings = get_settings(database_url="sqlite:///:memory:") + await create_all(engine=get_engine(settings=settings)) + return settings + + +@pytest_asyncio.fixture +async def async_session(default_settings: Settings): """Provide a clean asynchronous database session for async tests.""" - async with get_session() as async_session: + async with get_session(settings=default_settings) as async_session: yield async_session await dispose_database_runtime() diff --git a/tests/services/test_job_service.py b/tests/services/test_job_service.py index 7779c27..5a3274b 100644 --- a/tests/services/test_job_service.py +++ b/tests/services/test_job_service.py @@ -2,15 +2,18 @@ from uuid import uuid4 import pytest +from transcription.config import Settings +from transcription.db.runtime import get_session_factory from transcription.models import Job from transcription.services.jobs import JobService from transcription.services.jobs import JobStatus @pytest.fixture -def job_service(): +def job_service(default_settings: Settings) -> JobService: """Provide a JobService instance for testing.""" - return JobService() + session_factory = get_session_factory(settings=default_settings) + return JobService(session_factory=session_factory) class TestJobService: