generated from john/python-template
138 lines
4.2 KiB
Python
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]
|