fix: EventSession.commit() discards eager-loaded relations on watched models

This commit is contained in:
2026-08-28 13:28:22 -04:00
parent 6dafd40277
commit 654347126d
2 changed files with 126 additions and 13 deletions
+36 -7
View File
@@ -8,6 +8,7 @@ from typing import Any
from sqlalchemy import event, select, tuple_ from sqlalchemy import event, select, tuple_
from sqlalchemy import inspect as sa_inspect from sqlalchemy import inspect as sa_inspect
from sqlalchemy.ext.asyncio import AsyncSession 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 sqlalchemy.orm.attributes import set_committed_value as _sa_set_committed_value
from ..logger import get_logger from ..logger import get_logger
@@ -189,17 +190,44 @@ async def _invoke_callback(
await result 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( 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: ) -> 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_cols = sa_inspect(model, raiseerr=True).primary_key
pk_tuples = [sa_inspect(obj).key[1] for obj in objs]
where = ( where = (
pk_cols[0].in_([pk[0] for pk in pk_tuples]) pk_cols[0].in_([pk[0] for pk in pk_tuples])
if len(pk_cols) == 1 if len(pk_cols) == 1
else tuple_(*pk_cols).in_(pk_tuples) else tuple_(*pk_cols).in_(pk_tuples)
) )
q = select(model).where(where).execution_options(populate_existing=True) 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) await session.execute(q)
@@ -207,6 +235,7 @@ class EventSession(AsyncSession):
"""AsyncSession subclass that dispatches lifecycle callbacks after commit.""" """AsyncSession subclass that dispatches lifecycle callbacks after commit."""
async def commit(self) -> None: async def commit(self) -> None:
preloaded = _snapshot_loaded_relationships(self)
await super().commit() await super().commit()
creates: list[Any] = self.info.pop(_SESSION_CREATES, []) creates: list[Any] = self.info.pop(_SESSION_CREATES, [])
@@ -249,25 +278,25 @@ class EventSession(AsyncSession):
# session.get() per object. # session.get() per object.
create_items: list[Any] = [] create_items: list[Any] = []
update_items: list[tuple[Any, dict[str, dict[str, 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: for obj in creates:
state = sa_inspect(obj, raiseerr=False) state = sa_inspect(obj, raiseerr=False)
if state is None or state.detached or state.transient: # pragma: no cover if state is None or state.detached or state.transient: # pragma: no cover
continue continue
create_items.append(obj) 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(): for obj, changes in field_changes.values():
state = sa_inspect(obj, raiseerr=False) state = sa_inspect(obj, raiseerr=False)
if state is None or state.detached or state.transient: # pragma: no cover if state is None or state.detached or state.transient: # pragma: no cover
continue continue
update_items.append((obj, changes)) 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: try:
await _batch_reload(self, model, pk_tuples) await _batch_reload(self, model, objs, preloaded)
except Exception as exc: except Exception as exc:
_logger.error(_CALLBACK_ERROR_MSG, exc_info=exc) _logger.error(_CALLBACK_ERROR_MSG, exc_info=exc)
+90 -6
View File
@@ -6,9 +6,16 @@ from types import SimpleNamespace
from unittest.mock import patch from unittest.mock import patch
import pytest 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.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 import fastapi_toolsets.models.watched as _watched_module
from fastapi_toolsets.models import ( 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}) _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): class WatchAllModel(MixinBase, UUIDMixin):
"""Model without __watched_fields__ — watches all mapped fields by default.""" """Model without __watched_fields__ — watches all mapped fields by default."""
@@ -355,6 +383,62 @@ async def mixin_session_maker():
await engine.dispose() 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: class TestUUIDMixin:
@pytest.mark.anyio @pytest.mark.anyio
async def test_uuid_generated_by_db(self, mixin_session): async def test_uuid_generated_by_db(self, mixin_session):
@@ -1013,10 +1097,10 @@ class TestEventCallbacks:
real_batch_reload = _watched_module._batch_reload real_batch_reload = _watched_module._batch_reload
async def racing_batch_reload(session, model, pk_tuples): async def racing_batch_reload(session, model, objs, preloaded):
if any(pk[0] == doomed_id for pk in pk_tuples): if any(getattr(o, "id", None) == doomed_id for o in objs):
await kill_doomed_row_once() 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 # Patch the batched reload EventSession.commit() uses to pick up
# server defaults, so this test still exercises the race. # server defaults, so this test still exercises the race.
@@ -1039,7 +1123,7 @@ class TestEventCallbacks:
obj = WatchedModel(status="active", other="x") obj = WatchedModel(status="active", other="x")
mixin_session.add(obj) 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") raise RuntimeError("reload failed")
with ( with (