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) ]