Source code for fsh_lib.resource_registry

from __future__ import annotations

from dataclasses import dataclass
from decimal import Decimal
from typing import TYPE_CHECKING, Any, Literal, Protocol, assert_never

from pydantic import BaseModel
from sqlalchemy import (
    Text,
    cast,
    column,
    false,
    func,
    literal,
    select,
    true,
    union_all,
)
from sqlalchemy import (
    inspect as sa_inspect,
)

from fsh_lib.actions import ActionRef, ActionSpec, available_actions
from fsh_lib.filters import FilterOperator, trigram_clause
from fsh_lib.numeric import (
    DecimalString,  # noqa: TC001 -- pydantic field, needs runtime resolution
)
from fsh_lib.values_table import values_table

if TYPE_CHECKING:
    import enum as _enum_mod
    from collections.abc import (
        Awaitable,
        Callable,
        Collection,
        Iterable,
        Mapping,
        Sequence,
    )

    from sqlalchemy.ext.asyncio import AsyncSession
    from sqlalchemy.orm import DeclarativeBase
    from sqlalchemy.sql import ColumnElement, Select

    from fsh_lib.filter_values import FilterValuesRequest


[docs] class RowFilterFactory(Protocol): def __call__( self, *, action: str, resource_type: str, columns: Mapping[str, Any], ) -> Awaitable[ColumnElement[bool] | None]: ...
async def default_row_filter_factory( action: str, # noqa: ARG001 resource_type: str, # noqa: ARG001 columns: Mapping[str, Any], # noqa: ARG001 ) -> ColumnElement[bool] | None: return None type DefaultRepSerializer = Any @dataclass(frozen=True) class _ChoiceRow: value: str
[docs] @dataclass(frozen=True) class Enum: name: str enum_class: type[_enum_mod.Enum] operators: tuple[FilterOperator, ...] = ( FilterOperator.EQ, FilterOperator.IN, ) kind: Literal["enum"] = "enum" def members(self) -> list[_ChoiceRow]: return [ _ChoiceRow(value=str(member.value)) for member in self.enum_class ]
[docs] @dataclass(frozen=True) class Ref: name: str target: str operators: tuple[FilterOperator, ...] = ( FilterOperator.EQ, FilterOperator.IN, ) kind: Literal["ref"] = "ref"
[docs] @dataclass(frozen=True) class LiteralField: name: str type: str operators: tuple[FilterOperator, ...] = ( FilterOperator.EQ, FilterOperator.GT, FilterOperator.GTE, FilterOperator.LT, FilterOperator.LTE, ) kind: Literal["literal"] = "literal"
[docs] @dataclass(frozen=True) class Bool: name: str operators: tuple[FilterOperator, ...] = (FilterOperator.EQ,) kind: Literal["bool"] = "bool" def members(self) -> list[_ChoiceRow]: return [ _ChoiceRow(value="true"), _ChoiceRow(value="false"), ]
FilterField = Enum | Ref | LiteralField | Bool
[docs] class Result(BaseModel): field: str value: str item: Any | None = None score: DecimalString
[docs] class ValuesPage(BaseModel): results: list[Result]
[docs] @dataclass(frozen=True) class ResourceEntry: resource: str model: type pk: str default_rep_serializer: DefaultRepSerializer default_rep_class: type fields: tuple[FilterField, ...] = () search_columns: tuple[str, ...] = () opa_rep_class: type | None = None object_actions: tuple[ActionSpec, ...] = () collection_actions: tuple[ActionSpec, ...] = () @property def pk_column(self) -> ColumnElement[Any]: return getattr(self.model, self.pk) @property def pk_column_text(self) -> ColumnElement[str]: return cast(self.pk_column, Text) @property def label_column(self) -> ColumnElement[Any]: return ( getattr(self.model, self.search_columns[0]) if self.search_columns else self.pk_column_text )
[docs] class ResourceRegistry[Slug: str = str]: def __init__(self, entries: dict[Slug, ResourceEntry]) -> None: self._entries: dict[str, ResourceEntry] = { str(slug): entry for slug, entry in entries.items() } def __getitem__(self, slug: Slug | str) -> _BoundResource[Slug]: return _BoundResource( resource=self._entries[str(slug)], registry=self, # type: ignore[arg-type] ) def get_by_model( self, model: type[DeclarativeBase] | DeclarativeBase ) -> _BoundResource[Slug]: for entry in self._entries.values(): if model is entry.model or isinstance(model, entry.model): return _BoundResource( resource=entry, registry=self, # type: ignore[arg-type] ) raise KeyError(f"No resource registered for model {model}") def permission_codes(self) -> list[str]: codes: set[str] = {"*"} for slug, entry in self._entries.items(): codes.add(f"{slug}:*") for spec in (*entry.object_actions, *entry.collection_actions): codes.add(f"{slug}:{spec.name}") return sorted(codes) async def hydrate_refs( self, resource: Slug, ids: Iterable[Any], db: AsyncSession, auth: Any, # noqa: ANN401 ) -> dict[str, Any]: entry = self._entries[resource] stmt = select(entry.model).where(entry.pk_column.in_(list(ids))) load_options: Callable = getattr( entry.default_rep_class, "get_load_options", list, ) stmt = stmt.options(*load_options()) rows = (await db.execute(stmt)).scalars().all() return { getattr( row, entry.pk, ): await entry.default_rep_serializer(row, auth, db) for row in rows }
class _BoundResource[Slug: str = str]: resource: ResourceEntry registry: ResourceRegistry filter_fields: dict[str, FilterField] def __init__( self, resource: ResourceEntry, registry: ResourceRegistry, ) -> None: self.resource = resource self.registry = registry self.filter_fields = { field.name: field for field in self.resource.fields } async def actions[T: BaseModel = ActionRef]( self, *, auth: Any, # noqa: ANN401 obj: Any = None, # noqa: ANN401 ref_cls: type[T] = ActionRef, # type: ignore[assignment] allowed: Collection[str] | None = None, ) -> list[T]: action_specs = ( self.resource.object_actions if obj is not None else self.resource.collection_actions ) return await available_actions( resource=obj, auth=auth, action_specs=action_specs, ref_cls=ref_cls, allowed=allowed, ) async def view_filter( self, *, row_filter: RowFilterFactory = default_row_filter_factory, ) -> ColumnElement[bool]: clause = await self._compile_slug_filter( resource=self.resource, row_filter=row_filter, ) return true() if clause is None else clause async def values( self, request: FilterValuesRequest, db: AsyncSession, auth: Any, # noqa: ANN401 row_filter: RowFilterFactory = default_row_filter_factory, ) -> ValuesPage: fields = request.filter_fields or list(self.resource.search_columns) if not fields: return ValuesPage(results=[]) filters = await self._compile_row_filters( filter_fields=fields, row_filter=row_filter, ) if request.mode == "ids": id_subqueries = [ self._construct_id_subquery( entry=self.resource, field_name=field, ids=tuple(request.ids[field]), filters=filters, ) for field in fields if field in request.ids ] statement = union_all(*id_subqueries).order_by( column("value").asc(), ) elif request.mode == "trigram": trigram_subqueries = [ self._construct_trigram_subquery( entry=self.resource, field_name=field, query=request.search.strip(), filters=filters, ) for field in fields ] order_by: list = [ column("value").asc(), ] if request.search: order_by.insert(0, column("score").desc()) statement = union_all(*trigram_subqueries).order_by(*order_by) else: assert_never(request.mode) rows = (await db.execute(statement.limit(request.computed_limit))).all() ref_items = await self._hydrate_ref_items( filter_fields=fields, rows=rows, db=db, auth=auth, ) return ValuesPage( results=[ Result( field=row.field, value=row.value, item=ref_items.get(row.field, {}).get(str(row.value)), score=Decimal(row.score), ) for row in rows ], ) async def _hydrate_ref_items( self, *, filter_fields: list[str], rows: Sequence[Any], db: AsyncSession, auth: Any, # noqa: ANN401 ) -> dict[str, dict[str, dict[str, Any]]]: ref_fields = { filter_field: spec for filter_field in filter_fields if isinstance((spec := self.filter_fields.get(filter_field)), Ref) } field_id: dict[str, set] = {} for field, spec in ref_fields.items(): ids = {row.value for row in rows if row.field == field} field_id.setdefault(spec.target, set()).update(ids) items_by_field: dict[str, dict[str, Any]] = {} for target, ids in field_id.items(): hydrated = await self.registry.hydrate_refs( resource=target, ids=ids, db=db, auth=auth, ) items_by_field.update( { field: {str(pk): item for pk, item in hydrated.items()} for field, spec in ref_fields.items() if spec.target == target } ) return items_by_field async def _compile_slug_filter( self, resource: ResourceEntry, row_filter: RowFilterFactory, ) -> ColumnElement[bool] | None: columns: dict[str, Any] = {"id": resource.pk_column} representation = resource.opa_rep_class if representation is not None: mapped = sa_inspect(resource.model).columns columns |= { name: getattr(resource.model, name) for name in representation.model_fields if name in mapped and name not in ("id", "type") } return await row_filter( action=f"{resource.resource}:view", resource_type=resource.resource, columns=columns, ) async def _compile_row_filters( self, filter_fields: list[str], row_filter: RowFilterFactory = default_row_filter_factory, ) -> dict[str, ColumnElement[bool]]: targets: dict[str, ResourceEntry] = {} for filter_field in filter_fields: spec = self.filter_fields.get(filter_field) if isinstance(spec, (Enum, Bool)): continue target = ( self.registry[spec.target].resource if isinstance(spec, Ref) else self.resource ) targets[target.resource] = target clauses = {} for target in targets.values(): clause = await self._compile_slug_filter(target, row_filter) if clause is not None: clauses[target.resource] = clause return clauses def _construct_trigram_subquery( self, entry: ResourceEntry, field_name: str, query: str, filters: Mapping[str, ColumnElement[bool]], ) -> Select[Any]: spec = self.filter_fields.get(field_name) if isinstance(spec, (Enum, Bool)): members = values_table( dataclass_type=_ChoiceRow, instances=spec.members(), name=f"choice_{field_name}", ) score = ( func.similarity(members.c.value, query) if query else literal(0.0) ) stmt = select( literal(field_name).label("field"), members.c.value.label("value"), score.label("score"), ) if query: stmt = stmt.where(trigram_clause(members.c.value, query)) return stmt if isinstance(spec, Ref): target = self.registry[spec.target] stmt = select( literal(spec.name).label("field"), target.resource.pk_column_text.label("value"), func.similarity(target.resource.label_column, query).label( "score" ), ) if query: stmt = stmt.where( trigram_clause(target.resource.label_column, query) ) if (clause := filters.get(spec.target)) is not None: stmt = stmt.where(clause) return stmt if spec is None or isinstance(spec, LiteralField): column_attr = getattr(entry.model, field_name) score = ( func.similarity(column_attr, query) if query else literal(0.0) ) stmt = ( select( literal(field_name).label("field"), column_attr.label("value"), score.label("score"), ) .distinct() .where(column_attr.isnot(None), column_attr != "") ) if query: stmt = stmt.where(trigram_clause(column_attr, query)) if (clause := filters.get(self.resource.resource)) is not None: stmt = stmt.where(clause) return stmt assert_never(spec) def _construct_id_subquery( self, entry: ResourceEntry, field_name: str, ids: tuple[str, ...], filters: Mapping[str, ColumnElement[bool]], ) -> Select[Any]: spec = self.filter_fields.get(field_name) id_set = set(ids) _always_include_sentinel = literal(101) if isinstance(spec, (Enum, Bool)): matched = [row for row in spec.members() if row.value in id_set] if not matched: return select( literal(field_name).label("field"), literal("").label("value"), literal(0.0).label("score"), ).where(false()) members = values_table( dataclass_type=_ChoiceRow, instances=matched, name=f"choice_{field_name}", ) return select( literal(field_name).label("field"), members.c.value.label("value"), _always_include_sentinel.label("score"), ) if isinstance(spec, Ref): target = self.registry[spec.target] stmt = select( literal(spec.name).label("field"), target.resource.pk_column_text.label("value"), _always_include_sentinel.label("score"), ).where(target.resource.pk_column_text.in_(list(ids))) if (clause := filters.get(spec.target)) is not None: stmt = stmt.where(clause) return stmt if spec is None or isinstance(spec, LiteralField): column_attr = getattr(entry.model, field_name) stmt = ( select( literal(field_name).label("field"), column_attr.label("value"), _always_include_sentinel.label("score"), ) .distinct() .where(column_attr.in_(list(ids))) ) if (clause := filters.get(self.resource.resource)) is not None: stmt = stmt.where(clause) return stmt assert_never(spec)