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)