Compare commits

...
2 Commits
Author SHA1 Message Date
d3vyce ef269833b9 Version 5.1.2 2026-08-31 16:51:38 -04:00
d3vyce 610b3e1ab4 revert: drop populate_existing reload in create()/update() 2026-08-31 16:51:04 -04:00
5 changed files with 41 additions and 43 deletions
+1 -1
View File
@@ -1,6 +1,6 @@
[project] [project]
name = "fastapi-toolsets" name = "fastapi-toolsets"
version = "5.1.1" version = "5.1.2"
description = "Production-ready utilities for FastAPI applications" description = "Production-ready utilities for FastAPI applications"
readme = "README.md" readme = "README.md"
license = "MIT" license = "MIT"
+1 -1
View File
@@ -24,4 +24,4 @@ Example usage:
return Response(data={"user": user.username}, message="Success") return Response(data={"user": user.username}, message="Success")
""" """
__version__ = "5.1.1" __version__ = "5.1.2"
+12 -38
View File
@@ -270,32 +270,17 @@ class AsyncCrud(Generic[ModelType]):
return cls.default_load_options return cls.default_load_options
@classmethod @classmethod
def _capture_pk_values( async def _reload_with_options(
cls: type[Self], instance: DeclarativeBase cls: type[Self], session: AsyncSession, instance: DeclarativeBase
) -> dict[str, Any]:
"""Capture PK values off instance — call before commit expires attributes."""
return {
cast(str, col.key): getattr(instance, cast(str, col.key))
for col in cls.model.__mapper__.primary_key
}
@classmethod
async def _reload_with_options_by_pk(
cls: type[Self], session: AsyncSession, pk_values: dict[str, Any]
) -> ModelType: ) -> ModelType:
"""Re-query by previously captured PK values, with default_load_options applied.""" """Re-query instance by PK with default_load_options applied."""
# Only called when cls.default_load_options is set (see call sites). mapper = cls.model.__mapper__
pk_filters = [ pk_filters = [
getattr(cls.model, key) == value for key, value in pk_values.items() getattr(cls.model, cast(str, col.key))
== getattr(instance, cast(str, col.key))
for col in mapper.primary_key
] ]
q = select(cls.model).where(and_(*pk_filters)) return await cls.get(session, filters=pk_filters)
q = q.execution_options(populate_existing=True)
q = q.options(*cast(Sequence[ExecutableOption], cls.default_load_options))
result = await session.execute(q)
item = result.unique().scalar_one_or_none()
if item is None: # pragma: no cover — row was just flushed in this transaction
raise NotFoundError()
return cast(ModelType, item)
@classmethod @classmethod
async def _resolve_m2m( async def _resolve_m2m(
@@ -849,14 +834,9 @@ class AsyncCrud(Generic[ModelType]):
setattr(db_model, rel_attr, related_instances) setattr(db_model, rel_attr, related_instances)
session.add(db_model) session.add(db_model)
pk_values: dict[str, Any] | None = None
if cls.default_load_options:
await session.flush()
pk_values = cls._capture_pk_values(db_model)
if pk_values is not None:
db_model = await cls._reload_with_options_by_pk(session, pk_values)
else:
await session.refresh(db_model) await session.refresh(db_model)
if cls.default_load_options:
db_model = await cls._reload_with_options(session, db_model)
result = cast(ModelType, db_model) result = cast(ModelType, db_model)
if schema: if schema:
return Response(data=schema.model_validate(result)) return Response(data=schema.model_validate(result))
@@ -1233,15 +1213,9 @@ class AsyncCrud(Generic[ModelType]):
m2m_resolved = await cls._resolve_m2m(session, obj, only_set=True) m2m_resolved = await cls._resolve_m2m(session, obj, only_set=True)
for rel_attr, related_instances in m2m_resolved.items(): for rel_attr, related_instances in m2m_resolved.items():
setattr(db_model, rel_attr, related_instances) setattr(db_model, rel_attr, related_instances)
pk_values: dict[str, Any] | None = None
if cls.default_load_options:
await session.flush()
pk_values = cls._capture_pk_values(db_model)
if pk_values is not None:
db_model = await cls._reload_with_options_by_pk(session, pk_values)
else:
await session.refresh(db_model) await session.refresh(db_model)
if cls.default_load_options:
db_model = await cls._reload_with_options(session, db_model)
if schema: if schema:
return Response(data=schema.model_validate(db_model)) return Response(data=schema.model_validate(db_model))
return db_model return db_model
+24
View File
@@ -466,6 +466,30 @@ class TestDefaultLoadOptionsIntegration:
assert updated.role is not None assert updated.role is not None
assert updated.role.name == "admin" assert updated.role.name == "admin"
@pytest.mark.anyio
async def test_create_does_not_expire_already_loaded_relationships(
self, db_session: AsyncSession
):
"""create()'s reload must not blow away loaded state on related objects."""
UserWithDefaultLoad = CrudFactory(
User, default_load_options=[selectinload(User.role)]
)
role = await RoleCrud.create(db_session, RoleCreate(name="admin"))
role = await RoleCrud.get(
db_session,
filters=[Role.id == role.id],
load_options=[selectinload(Role.users)],
)
assert role.users == []
await UserWithDefaultLoad.create(
db_session,
UserCreate(username="alice", email="alice@test.com", role_id=role.id),
)
# must not trigger a lazy load
assert role.users == []
@pytest.mark.anyio @pytest.mark.anyio
async def test_load_options_overrides_default_load_options( async def test_load_options_overrides_default_load_options(
self, db_session: AsyncSession self, db_session: AsyncSession
Generated
+1 -1
View File
@@ -315,7 +315,7 @@ wheels = [
[[package]] [[package]]
name = "fastapi-toolsets" name = "fastapi-toolsets"
version = "5.1.1" version = "5.1.2"
source = { editable = "." } source = { editable = "." }
dependencies = [ dependencies = [
{ name = "asyncpg" }, { name = "asyncpg" },