diff --git a/src/fastapi_toolsets/models/watched.py b/src/fastapi_toolsets/models/watched.py index fb6ef07..90b17b1 100644 --- a/src/fastapi_toolsets/models/watched.py +++ b/src/fastapi_toolsets/models/watched.py @@ -8,6 +8,7 @@ from typing import Any from sqlalchemy import event, select, tuple_ from sqlalchemy import inspect as sa_inspect from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy.orm import selectinload from sqlalchemy.orm.attributes import set_committed_value as _sa_set_committed_value from ..logger import get_logger @@ -189,17 +190,44 @@ async def _invoke_callback( await result +def _loaded_relationships(obj: Any) -> set[str]: + """Relationship keys currently loaded on *obj*.""" + state = sa_inspect(obj) + unloaded = state.unloaded + return { + rel.key + for rel in state.mapper.relationships + if rel.key not in unloaded and rel.lazy not in ("dynamic", "write_only") + } + + +def _snapshot_loaded_relationships(session: Any) -> dict[int, set[str]]: + """Record loaded relationships for the tracked objects, keyed by ``id``.""" + objs = list(session.info.get(_SESSION_CREATES, [])) + objs += [obj for obj, _ in session.info.get(_SESSION_UPDATES, {}).values()] + return {id(obj): _loaded_relationships(obj) for obj in objs} + + async def _batch_reload( - session: AsyncSession, model: type, pk_tuples: list[tuple[Any, ...]] + session: AsyncSession, + model: type, + objs: list[Any], + preloaded: dict[int, set[str]], ) -> None: - """Re-populate all rows of *model* identified by *pk_tuples* in one round trip.""" + """Re-populate all rows of *model* in one round trip.""" pk_cols = sa_inspect(model, raiseerr=True).primary_key + pk_tuples = [sa_inspect(obj).key[1] for obj in objs] where = ( pk_cols[0].in_([pk[0] for pk in pk_tuples]) if len(pk_cols) == 1 else tuple_(*pk_cols).in_(pk_tuples) ) q = select(model).where(where).execution_options(populate_existing=True) + loaded: set[str] = set() + for obj in objs: + loaded |= preloaded.get(id(obj), set()) + if loaded: + q = q.options(*(selectinload(getattr(model, key)) for key in loaded)) await session.execute(q) @@ -207,6 +235,7 @@ class EventSession(AsyncSession): """AsyncSession subclass that dispatches lifecycle callbacks after commit.""" async def commit(self) -> None: + preloaded = _snapshot_loaded_relationships(self) await super().commit() creates: list[Any] = self.info.pop(_SESSION_CREATES, []) @@ -249,25 +278,25 @@ class EventSession(AsyncSession): # session.get() per object. create_items: list[Any] = [] update_items: list[tuple[Any, dict[str, dict[str, Any]]]] = [] - pk_by_type: dict[type, list[tuple[Any, ...]]] = {} + objs_by_type: dict[type, list[Any]] = {} for obj in creates: state = sa_inspect(obj, raiseerr=False) if state is None or state.detached or state.transient: # pragma: no cover continue create_items.append(obj) - pk_by_type.setdefault(type(obj), []).append(state.key[1]) + objs_by_type.setdefault(type(obj), []).append(obj) for obj, changes in field_changes.values(): state = sa_inspect(obj, raiseerr=False) if state is None or state.detached or state.transient: # pragma: no cover continue update_items.append((obj, changes)) - pk_by_type.setdefault(type(obj), []).append(state.key[1]) + objs_by_type.setdefault(type(obj), []).append(obj) - for model, pk_tuples in pk_by_type.items(): + for model, objs in objs_by_type.items(): try: - await _batch_reload(self, model, pk_tuples) + await _batch_reload(self, model, objs, preloaded) except Exception as exc: _logger.error(_CALLBACK_ERROR_MSG, exc_info=exc) diff --git a/tests/test_models.py b/tests/test_models.py index 42a5062..a220459 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -6,9 +6,16 @@ from types import SimpleNamespace from unittest.mock import patch import pytest -from sqlalchemy import String +from sqlalchemy import ForeignKey, String, select +from sqlalchemy import inspect as sa_inspect from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine -from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column +from sqlalchemy.orm import ( + DeclarativeBase, + Mapped, + mapped_column, + relationship, + selectinload, +) import fastapi_toolsets.models.watched as _watched_module from fastapi_toolsets.models import ( @@ -107,6 +114,27 @@ async def _watched_on_update(obj, event_type, changes): _test_events.append({"event": "update", "obj_id": obj.id, "changes": changes}) +class RelTarget(MixinBase, UUIDMixin): + __tablename__ = "mixin_rel_targets" + + name: Mapped[str] = mapped_column(String(50)) + + +class RelOwner(MixinBase, UUIDMixin): + """Watched model with a relationship, to check eager loads survive commit.""" + + __tablename__ = "mixin_rel_owners" + + title: Mapped[str] = mapped_column(String(50)) + target_id: Mapped[uuid.UUID] = mapped_column(ForeignKey("mixin_rel_targets.id")) + target: Mapped[RelTarget] = relationship() + + +@listens_for(RelOwner, [ModelEvent.CREATE, ModelEvent.UPDATE]) +async def _rel_owner_handler(obj, event_type, changes): + _test_events.append({"event": event_type.value, "obj_id": obj.id}) + + class WatchAllModel(MixinBase, UUIDMixin): """Model without __watched_fields__ — watches all mapped fields by default.""" @@ -355,6 +383,62 @@ async def mixin_session_maker(): await engine.dispose() +class TestEventSessionPreservesEagerLoads: + """EventSession.commit() must not discard relations an eager load populated.""" + + async def _seed_eager(self, session): + target = RelTarget(name="t") + session.add(target) + await session.flush() + owner = RelOwner(title="o", target_id=target.id) + session.add(owner) + await session.flush() + loaded = ( + await session.execute( + select(RelOwner) + .where(RelOwner.id == owner.id) + .options(selectinload(RelOwner.target)) + ) + ).scalar_one() + assert "target" not in sa_inspect(loaded).unloaded + return loaded + + @pytest.mark.anyio + async def test_eager_load_survives_commit(self, mixin_session): + """expire_on_commit=False: the reload must not expire the relation.""" + owner = await self._seed_eager(mixin_session) + + await mixin_session.commit() + + assert "target" not in sa_inspect(owner).unloaded + assert owner.target.name == "t" + + @pytest.mark.anyio + async def test_eager_load_survives_commit_expire_on_commit( + self, mixin_session_expire + ): + """expire_on_commit=True: what was loaded must be recorded before the commit.""" + owner = await self._seed_eager(mixin_session_expire) + + await mixin_session_expire.commit() + + assert "target" not in sa_inspect(owner).unloaded + assert owner.target.name == "t" + + @pytest.mark.anyio + async def test_unloaded_relation_stays_unloaded(self, mixin_session): + """Only what was loaded is restored: the reload must not eager-load extra.""" + target = RelTarget(name="t") + mixin_session.add(target) + await mixin_session.flush() + owner = RelOwner(title="o", target_id=target.id) + mixin_session.add(owner) + + await mixin_session.commit() + + assert "target" in sa_inspect(owner).unloaded + + class TestUUIDMixin: @pytest.mark.anyio async def test_uuid_generated_by_db(self, mixin_session): @@ -1013,10 +1097,10 @@ class TestEventCallbacks: real_batch_reload = _watched_module._batch_reload - async def racing_batch_reload(session, model, pk_tuples): - if any(pk[0] == doomed_id for pk in pk_tuples): + async def racing_batch_reload(session, model, objs, preloaded): + if any(getattr(o, "id", None) == doomed_id for o in objs): await kill_doomed_row_once() - return await real_batch_reload(session, model, pk_tuples) + return await real_batch_reload(session, model, objs, preloaded) # Patch the batched reload EventSession.commit() uses to pick up # server defaults, so this test still exercises the race. @@ -1039,7 +1123,7 @@ class TestEventCallbacks: obj = WatchedModel(status="active", other="x") mixin_session.add(obj) - async def failing_batch_reload(session, model, pk_tuples): + async def failing_batch_reload(session, model, objs, preloaded): raise RuntimeError("reload failed") with (