Merge pull request #385 from d3vyce/384-eventsessioncommit-discards-eager-loaded-relations-on-watched-models

fix: EventSession.commit() discards eager-loaded relations on watched models
This commit is contained in:
d3vyce
2026-08-28 22:23:09 +02:00
committed by GitHub
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 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)
+90 -6
View File
@@ -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 (