Source code for fsh_lib.loading
from dataclasses import dataclass
from enum import StrEnum
from typing import TYPE_CHECKING
from sqlalchemy.orm import (
contains_eager,
joinedload,
selectinload,
subqueryload,
)
if TYPE_CHECKING:
from collections.abc import Callable, Iterable
from sqlalchemy import Select
from sqlalchemy.orm import InstrumentedAttribute, RelationshipProperty
from sqlalchemy.orm.strategy_options import _AbstractLoad
[docs]
class EagerStrategy(StrEnum):
SELECTIN = "selectin"
JOINED = "joined"
SUBQUERY = "subquery"
_HEAD_LOADER: dict[EagerStrategy, Callable[..., _AbstractLoad]] = {
EagerStrategy.SELECTIN: selectinload,
EagerStrategy.JOINED: joinedload,
EagerStrategy.SUBQUERY: subqueryload,
}
_CHAIN_METHOD: dict[EagerStrategy, str] = {
EagerStrategy.SELECTIN: "selectinload",
EagerStrategy.JOINED: "joinedload",
EagerStrategy.SUBQUERY: "subqueryload",
}
[docs]
@dataclass(frozen=True)
class EagerLoad:
attr: InstrumentedAttribute
strategy: EagerStrategy = EagerStrategy.SELECTIN
children: tuple[EagerLoad, ...] = ()
def apply_eager_loads(
stmt: Select,
loads: Iterable[EagerLoad],
joined: set[RelationshipProperty] | None = None,
) -> Select:
joined = joined or set()
for load in loads:
stmt = stmt.options(*_loaders(load, joined, head=True))
return stmt
def _loaders(
load: EagerLoad,
joined: set[RelationshipProperty],
*,
head: bool,
) -> list[_AbstractLoad]:
if head and load.attr.property in joined:
option = contains_eager(load.attr)
else:
option = _HEAD_LOADER[load.strategy](load.attr)
if not load.children:
return [option]
return [
extended
for child in load.children
for extended in _extend(option, child)
]
def _extend(option: _AbstractLoad, child: EagerLoad) -> list[_AbstractLoad]:
sub: _AbstractLoad = getattr(option, _CHAIN_METHOD[child.strategy])(
child.attr
)
if not child.children:
return [sub]
return [
extended
for grandchild in child.children
for extended in _extend(sub, grandchild)
]