Files
transcription/tests/services/test_service_base.py
T

138 lines
4.2 KiB
Python

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]