Files
transcription/tests/services/test_service_base.py
T
2026-08-20 16:14:44 -05:00

157 lines
4.8 KiB
Python

from __future__ import annotations
from contextlib import asynccontextmanager
from typing import cast
import pytest
from sqlalchemy.ext.asyncio import async_sessionmaker
from sqlmodel.ext.asyncio.session import AsyncSession
from transcription.config import Settings
from transcription.services.base import ServiceBase
class _TrackingSession:
def __init__(self) -> None:
self.commits = 0
self.flushes = 0
self.refreshed: list[object] = []
async def commit(self) -> None:
self.commits += 1
async def flush(self) -> None:
self.flushes += 1
async def refresh(self, obj: object) -> None:
self.refreshed.append(obj)
def _settings() -> Settings:
return Settings(_env_file=None, openrouter_api_key="test-key")
def _session_factory_stub() -> async_sessionmaker[AsyncSession]:
return cast(async_sessionmaker[AsyncSession], object())
def _session_stub() -> AsyncSession:
return cast(AsyncSession, object())
def test_initializes_with_defaults(monkeypatch):
settings = _settings()
session_factory = _session_factory_stub()
monkeypatch.setattr("transcription.services.base.get_settings", lambda: settings)
monkeypatch.setattr(
"transcription.services.base.resolve_session_factory",
lambda **_kwargs: session_factory,
)
service = ServiceBase()
assert service.settings is settings
assert service.session_factory is session_factory
def test_initializes_with_custom_session_factory():
settings = _settings()
session_factory = _session_factory_stub()
service = ServiceBase(settings=settings, session_factory=session_factory)
assert service.settings is settings
assert service.session_factory is session_factory
@pytest.mark.asyncio
async def test_session_scope_reuses_provided_session(monkeypatch):
captured: dict[str, object | None] = {}
provided_session = _session_stub()
yielded = _session_stub()
service = ServiceBase(settings=_settings(), session_factory=_session_factory_stub())
@asynccontextmanager
async def _fake_session_scope(*, session_factory=None, session=None):
captured["session_factory"] = session_factory
captured["session"] = session
yield yielded if session is None else session
monkeypatch.setattr("transcription.services.base.session_scope", _fake_session_scope)
async with service._session_scope(session=provided_session) as active:
assert active is provided_session
assert captured == {"session_factory": service.session_factory, "session": provided_session}
@pytest.mark.asyncio
async def test_session_scope_creates_owned_session_when_none_provided(monkeypatch):
captured: dict[str, object | None] = {}
owned_session = _session_stub()
service = ServiceBase(settings=_settings(), session_factory=_session_factory_stub())
@asynccontextmanager
async def _fake_session_scope(*, session_factory=None, session=None):
captured["session_factory"] = session_factory
captured["session"] = session
yield owned_session
monkeypatch.setattr("transcription.services.base.session_scope", _fake_session_scope)
async with service._session_scope() as active:
assert active is owned_session
assert captured == {"session_factory": service.session_factory, "session": None}
@pytest.mark.asyncio
async def test_session_scope_propagates_exceptions(monkeypatch):
service = ServiceBase(settings=_settings(), session_factory=_session_factory_stub())
@asynccontextmanager
async def _fake_session_scope(*, session_factory=None, session=None):
_ = (session_factory, session)
yield _session_stub()
monkeypatch.setattr("transcription.services.base.session_scope", _fake_session_scope)
with pytest.raises(RuntimeError, match="boom"):
async with service._session_scope():
raise RuntimeError("boom")
@pytest.mark.asyncio
async def test_finalize_commits_for_service_owned_session():
service = ServiceBase(settings=_settings(), session_factory=_session_factory_stub())
session = _TrackingSession()
refreshed = object()
await service._finalize(
session=cast(AsyncSession, session),
caller_session=None,
refresh=(refreshed,),
)
assert session.commits == 1
assert session.flushes == 0
assert session.refreshed == [refreshed]
@pytest.mark.asyncio
async def test_finalize_flushes_for_caller_owned_session():
service = ServiceBase(settings=_settings(), session_factory=_session_factory_stub())
session = _TrackingSession()
refreshed = object()
await service._finalize(
session=cast(AsyncSession, session),
caller_session=_session_stub(),
refresh=(refreshed,),
)
assert session.commits == 0
assert session.flushes == 1
assert session.refreshed == [refreshed]