from __future__ import annotations from contextlib import asynccontextmanager import pytest 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 test_initializes_with_defaults(monkeypatch): settings = _settings() session_factory = object() 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 = object() 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 = object() yielded = object() service = ServiceBase(settings=_settings(), session_factory=object()) @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 = object() service = ServiceBase(settings=_settings(), session_factory=object()) @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=object()) @asynccontextmanager async def _fake_session_scope(*, session_factory=None, session=None): _ = (session_factory, session) yield object() 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=object()) session = _TrackingSession() refreshed = object() await service._finalize(session=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=object()) session = _TrackingSession() refreshed = object() await service._finalize(session=session, caller_session=object(), refresh=(refreshed,)) assert session.commits == 0 assert session.flushes == 1 assert session.refreshed == [refreshed]