diff --git a/fastapi_admin_kit/admin/admin_database.py b/fastapi_admin_kit/admin/admin_database.py index 16aeacb..2d89067 100644 --- a/fastapi_admin_kit/admin/admin_database.py +++ b/fastapi_admin_kit/admin/admin_database.py @@ -6,6 +6,7 @@ import re from typing import Any +from fastapi_admin_kit.backends.sqlalchemy import SqlAlchemyDatabaseBackend from fastapi_admin_kit.config.database import DatabaseConfig logger = logging.getLogger(__name__) @@ -21,7 +22,11 @@ def _validate_identifier(name: str, kind: str = "table") -> str: class AdminDatabase: - """Handles database setup, table creation, and role seeding.""" + """Handles database setup, table creation, and role seeding. + + Delegates engine creation, table creation, and auto-migration to + :class:`SqlAlchemyDatabaseBackend`. + """ def __init__( self, @@ -32,11 +37,12 @@ def __init__( self.engine = engine self.base = base self.database_config = database_config + self._backend = SqlAlchemyDatabaseBackend( + admin_database=self, database_config=database_config + ) def _ensure_engine(self) -> Any: - """ - Create the async engine from ``database_config`` if no engine is set. - """ + """Create the async engine from ``database_config`` if no engine is set.""" if self.engine is None and self.database_config is not None: self.engine = self.database_config.create_engine() return self.engine @@ -108,9 +114,7 @@ def _auto_migrate_sync(self, metadata: Any) -> None: conn.execute(sql) def _auto_migrate(self, sync_conn: Any, metadata: Any) -> None: - """ - Add missing columns to existing tables (sync, called via run_sync). - """ + """Add missing columns to existing tables (sync, called via run_sync).""" from sqlalchemy import inspect as sa_inspect from sqlalchemy import text diff --git a/fastapi_admin_kit/admin/core.py b/fastapi_admin_kit/admin/core.py index db197eb..2a51811 100644 --- a/fastapi_admin_kit/admin/core.py +++ b/fastapi_admin_kit/admin/core.py @@ -721,6 +721,16 @@ def _wire_app_state(self, app: FastAPI) -> None: app.state.admin_jinja_env = state.jinja_env # Unified signing-key source for sessions, CSRF, and JWT (see AdminState). app.state.admin_secret_key = state.secret_key + # Multi-ORM backend: store adapter class for views to access + from fastapi_admin_kit.backends.sqlalchemy import ( + SqlAlchemyIntrospectionAdapter, + SqlAlchemyQueryAdapter, + SqlAlchemySessionAdapter, + ) + + app.state.admin_session_backend_class = SqlAlchemySessionAdapter + app.state.admin_query_adapter = SqlAlchemyQueryAdapter() + app.state.admin_introspection_adapter = SqlAlchemyIntrospectionAdapter() # Wire the password hasher to the User model from fastapi_admin_kit.auth.models import User diff --git a/fastapi_admin_kit/backends/__init__.py b/fastapi_admin_kit/backends/__init__.py new file mode 100644 index 0000000..0ed503b --- /dev/null +++ b/fastapi_admin_kit/backends/__init__.py @@ -0,0 +1,57 @@ +"""Backend protocols and SQLAlchemy adapters for multi-ORM support. + +Protocols:: + + from fastapi_admin_kit.backends import ( + IntrospectionBackend, + SessionBackend, + QueryBackend, + AuditBackend, + DatabaseBackend, + ) + +SQLAlchemy adapters:: + + from fastapi_admin_kit.backends import ( + SqlAlchemyIntrospectionAdapter, + SqlAlchemySessionAdapter, + SqlAlchemyQueryAdapter, + SqlAlchemyDatabaseBackend, + ) +""" + +from fastapi_admin_kit.backends.protocols import ( + AuditBackend, + ColumnMetaType, + DatabaseBackend, + IntrospectionBackend, + QueryBackend, + QueryType, + RelationMetaType, + SessionBackend, + SessionType, +) +from fastapi_admin_kit.backends.sqlalchemy import ( + SqlAlchemyDatabaseBackend, + SqlAlchemyIntrospectionAdapter, + SqlAlchemyQueryAdapter, + SqlAlchemySessionAdapter, +) + +__all__ = [ + # Protocols + "AuditBackend", + "ColumnMetaType", + "DatabaseBackend", + "IntrospectionBackend", + "QueryBackend", + "QueryType", + "RelationMetaType", + "SessionBackend", + "SessionType", + # SQLAlchemy adapters + "SqlAlchemyDatabaseBackend", + "SqlAlchemyIntrospectionAdapter", + "SqlAlchemyQueryAdapter", + "SqlAlchemySessionAdapter", +] diff --git a/fastapi_admin_kit/backends/protocols.py b/fastapi_admin_kit/backends/protocols.py new file mode 100644 index 0000000..e2d3b7f --- /dev/null +++ b/fastapi_admin_kit/backends/protocols.py @@ -0,0 +1,187 @@ +"""Protocol interfaces for multi-ORM backend support. + +These protocols define the contracts that all ORM backends (SQLAlchemy, +MongoDB/ODM, future) must implement. They are the seam that decouples +the rest of the codebase from any specific ORM. + +Use structural subtyping — any class that satisfies the protocol's +method signatures is a valid implementation, no inheritance required. +""" + +from __future__ import annotations + +from collections.abc import Sequence +from typing import Any, Protocol, TypeVar, runtime_checkable + +from fastapi_admin_kit.types import ColumnMeta, RelationMeta + +ModelT = TypeVar("ModelT") +ObjT = TypeVar("ObjT") + +# Type aliases — the rest of the codebase should reference these +# instead of SQLAlchemy-specific types. +QueryType = Any +SessionType = Any +ColumnMetaType = ColumnMeta +RelationMetaType = RelationMeta + + +@runtime_checkable +class IntrospectionBackend(Protocol): + """Model introspection: reflect columns, relationships, PKs, and abstractness.""" + + def inspect_model(self, model: type) -> tuple[list[ColumnMeta], list[RelationMeta]]: + """Inspect a model and return its column and relationship metadata.""" + ... + + def get_pk_field(self, model: type) -> str | tuple[str, ...] | None: + """Return the primary key field name(s) for a model.""" + ... + + def cast_pk_value(self, model: type, value: Any) -> Any: + """Cast a string PK value to the correct Python type for the model.""" + ... + + def is_abstract(self, model: type) -> bool: + """Return True if the model is abstract and should be skipped.""" + ... + + def get_relationship_names(self, model: type) -> set[str]: + """Return the set of relationship key names on a model.""" + ... + + def get_relationship(self, model: type, name: str) -> Any: + """Return a single relationship descriptor by name, or None.""" + ... + + def get_column_type_name(self, model: type, field_name: str) -> str | None: + """Return the SQLAlchemy type class name for a column, or None.""" + ... + + def get_column_attr(self, model: type, field_name: str) -> Any: + """Return the column attribute for a field name, or None.""" + ... + + def get_pk_columns(self, model: type) -> list[Any]: + """Return the primary key column(s) for a model.""" + ... + + +@runtime_checkable +class SessionBackend(Protocol): + """Data access: per-request session lifecycle.""" + + def get(self, model: type[ModelT], pk: Any) -> ModelT | None: + """Fetch a single object by primary key.""" + ... + + def add(self, obj: Any) -> None: + """Stage an object for insertion.""" + ... + + def flush(self) -> None: + """Flush pending changes to the DB without committing.""" + ... + + def delete(self, obj: Any) -> None: + """Mark an object for deletion.""" + ... + + def refresh(self, obj: Any, attributes: Sequence[str] | None = None) -> None: + """Re-read object attributes from the DB.""" + ... + + def execute(self, query: QueryType) -> Any: + """Execute a query object and return the result.""" + ... + + def commit(self) -> None: + """Persist all pending changes.""" + ... + + def rollback(self) -> None: + """Discard all pending changes.""" + ... + + +@runtime_checkable +class QueryBackend(Protocol): + """Chainable query building: select, filter, sort, join, paginate.""" + + def select(self, model: type[ModelT]) -> QueryType: + """Start a new query for the given model.""" + ... + + def where(self, query: QueryType, *conditions: Any) -> QueryType: + """Add WHERE conditions to a query.""" + ... + + def order_by(self, query: QueryType, *columns: Any) -> QueryType: + """Add ORDER BY clauses to a query.""" + ... + + def limit(self, query: QueryType, n: int) -> QueryType: + """Limit the result set to *n* rows.""" + ... + + def offset(self, query: QueryType, n: int) -> QueryType: + """Skip the first *n* rows of the result set.""" + ... + + def join(self, query: QueryType, related: type, on: Any | None = None) -> QueryType: + """Join a related model onto the query.""" + ... + + def distinct(self, query: QueryType) -> QueryType: + """Add DISTINCT to the query.""" + ... + + def count(self, query: QueryType) -> int: + """Execute the query and return the total row count.""" + ... + + def options(self, query: QueryType, *opts: Any) -> QueryType: + """Add eager-load options (joinedload, selectinload, etc.).""" + ... + + def ilike(self, column: Any, pattern: str) -> Any: + """Apply case-insensitive LIKE to a column, returning a boolean clause.""" + ... + + def or_(self, *clauses: Any) -> Any: + """Compose multiple boolean clauses with OR.""" + ... + + +@runtime_checkable +class AuditBackend(Protocol): + """Change tracking: attach listeners, snapshot, and diff objects.""" + + def attach_listeners(self, session_factory: Any, registry: dict[str, Any]) -> None: + """Register change-tracking listeners on the session factory.""" + ... + + def snapshot(self, obj: Any) -> dict[str, Any]: + """Capture a serialisable snapshot of the object's current state.""" + ... + + def compute_diff(self, before: dict[str, Any], after: dict[str, Any]) -> dict[str, Any]: + """Return a dict of {field: (old_value, new_value)} for changed fields.""" + ... + + +@runtime_checkable +class DatabaseBackend(Protocol): + """Connection lifecycle: create engine, run DDL, auto-migrate.""" + + def create_connection(self) -> Any: + """Create and return a new database connection or engine.""" + ... + + def create_tables(self, connection: Any, metadata: Any) -> None: + """Issue DDL to create all tables defined in *metadata*.""" + ... + + def auto_migrate(self, connection: Any, metadata: Any) -> None: + """Detect schema drift and apply migrations automatically.""" + ... diff --git a/fastapi_admin_kit/backends/sqlalchemy.py b/fastapi_admin_kit/backends/sqlalchemy.py new file mode 100644 index 0000000..8288406 --- /dev/null +++ b/fastapi_admin_kit/backends/sqlalchemy.py @@ -0,0 +1,465 @@ +"""SQLAlchemy backend adapters implementing the multi-ORM protocol interfaces. + +Contains: +- ``SqlAlchemyIntrospectionAdapter`` — model introspection (#23) +- ``SqlAlchemySessionAdapter`` — per-request session lifecycle (#24) +- ``SqlAlchemyQueryAdapter`` — chainable query building (#25) +- ``SqlAlchemyDatabaseBackend`` — connection lifecycle & DDL (#30) +""" + +from __future__ import annotations + +from collections.abc import Sequence +from typing import Any + +from sqlalchemy import inspect as sa_inspect + +from fastapi_admin_kit.types import ColumnMeta, RelationMeta + + +def _is_async_session(session: Any) -> bool: + """Return True if *session* is an SQLAlchemy async session.""" + from sqlalchemy.ext.asyncio import AsyncSession + + return isinstance(session, AsyncSession) + + +# --------------------------------------------------------------------------- +# #23 — Introspection Adapter +# --------------------------------------------------------------------------- + + +class SqlAlchemyIntrospectionAdapter: + """Reflects SQLAlchemy model metadata into ColumnMeta / RelationMeta. + + Implements :class:`IntrospectionBackend` via structural subtyping. + """ + + def inspect_model(self, model: type) -> tuple[list[ColumnMeta], list[RelationMeta]]: + """Inspect a SQLAlchemy model and return column + relationship metadata.""" + mapper = sa_inspect(model) + columns: list[ColumnMeta] = [] + relationships: list[RelationMeta] = [] + + is_sqlmodel = self._is_sqlmodel(model) + + for col in mapper.columns: + col_type = col.type + if is_sqlmodel: + col_type = self._resolve_sqlmodel_type(model, col.key, col.type) + columns.append( + ColumnMeta( + name=col.key, + type=col_type, + nullable=col.nullable, + primary_key=col.primary_key, + foreign_keys=list(col.foreign_keys), + default=col.default, + server_default=col.server_default, + index=col.index, + unique=col.unique, + ) + ) + + for rel in mapper.relationships: + try: + relationships.append( + RelationMeta( + name=rel.key, + direction=rel.direction.name, + target_model=rel.mapper.class_, + uselist=rel.uselist, + back_populates=rel.back_populates, + secondary=rel.secondary, + ) + ) + except Exception: + pass + + return columns, relationships + + def get_pk_field(self, model: type) -> str | tuple[str, ...] | None: + """Return the primary key field name(s) for a model.""" + mapper = sa_inspect(model) + pk_cols = mapper.primary_key + if not pk_cols: + return None + if len(pk_cols) == 1: + return pk_cols[0].key + return tuple(col.key for col in pk_cols) + + def cast_pk_value(self, model: type, value: Any) -> Any: + """Cast a string PK value to the correct Python type for the model.""" + if value is None: + return None + mapper = sa_inspect(model) + pk_cols = mapper.primary_key + if not pk_cols or len(pk_cols) != 1: + return value + pk_col = pk_cols[0] + from sqlalchemy import BigInteger, Integer + from sqlalchemy.dialects.postgresql import UUID as PG_UUID + from sqlalchemy.types import Uuid + + col_type = type(pk_col.type) + if col_type in (Integer, BigInteger): + return int(value) + if col_type in (PG_UUID, Uuid): + from uuid import UUID + + return UUID(str(value)) + return value + + def is_abstract(self, model: type) -> bool: + """Return True if the model is abstract and should be skipped.""" + return getattr(model, "__abstract__", False) + + def get_relationship_names(self, model: type) -> set[str]: + """Return the set of relationship key names on a model.""" + mapper = sa_inspect(model) + return {r.key for r in mapper.relationships} + + def get_relationship(self, model: type, name: str) -> Any: + """Return a single relationship descriptor by name, or None.""" + mapper = sa_inspect(model) + return mapper.relationships.get(name) + + def get_column_type_name(self, model: type, field_name: str) -> str | None: + """Return the SQLAlchemy type class name for a column, or None.""" + mapper = sa_inspect(model) + for prop in mapper.column_attrs: + if prop.key == field_name: + col = prop.columns[0] if prop.columns else None + if col is not None: + return col.type.__class__.__name__ + return None + + def get_column_attr(self, model: type, field_name: str) -> Any: + """Return the column attribute for a field name, or None.""" + mapper = sa_inspect(model) + for prop in mapper.column_attrs: + if prop.key == field_name: + col = prop.columns[0] if prop.columns else None + return col + return None + + def get_pk_columns(self, model: type) -> list[Any]: + """Return the primary key column(s) for a model.""" + mapper = sa_inspect(model) + return list(mapper.primary_key) + + # -- internal helpers --------------------------------------------------- + + def _is_sqlmodel(self, model: type) -> bool: + try: + from sqlmodel import SQLModel + + return isinstance(model, type) and issubclass(model, SQLModel) + except ImportError: + return False + + def _resolve_sqlmodel_type(self, model: type, field_name: str, default_type: Any) -> Any: + try: + from sqlmodel import SQLModel + + if not (isinstance(model, type) and issubclass(model, SQLModel)): + return default_type + + sqlmodel_fields = getattr(model, "__sqlmodel_fields__", {}) + if field_name not in sqlmodel_fields: + return default_type + + field_info = sqlmodel_fields[field_name] + annotation = getattr(field_info, "annotation", None) + if annotation is None: + return default_type + + import sqlalchemy as sa + + type_map = { + int: sa.Integer, + str: sa.String, + float: sa.Float, + bool: sa.Boolean, + } + + origin = getattr(annotation, "__origin__", None) + if origin is not None: + args = getattr(annotation, "__args__", ()) + if args: + inner = args[0] + if inner in type_map: + return type_map[inner] + + if annotation in type_map: + return type_map[annotation] + return default_type + except Exception: + return default_type + + +# --------------------------------------------------------------------------- +# #24 — Session Adapter +# --------------------------------------------------------------------------- + + +class SqlAlchemySessionAdapter: + """Wraps an ``AsyncSession`` (or sync ``Session``) to implement + :class:`SessionBackend`. + + When wrapping an ``AsyncSession``, all methods that talk to the DB + return awaitable coroutines so that existing ``await session.flush()`` + call-sites continue to work. + """ + + def __init__(self, session: Any) -> None: + self._session = session + self._is_async = hasattr(session, "__await__") or _is_async_session(session) + + @property + def session(self) -> Any: + return self._session + + def _maybe_async(self, coro: Any) -> Any: + """If the underlying session is async and we're in an async context, + return the coroutine so the caller can ``await`` it. + Otherwise run it synchronously and return the result.""" + if self._is_async and hasattr(coro, "__await__"): + return coro + if hasattr(coro, "__await__"): + import asyncio + + try: + loop = asyncio.get_running_loop() + except RuntimeError: + loop = None + if loop and loop.is_running(): + return coro + return loop.run_until_complete(coro) if loop else coro + return coro + + def get(self, model: type, pk: Any) -> Any | None: + """Fetch a single object by primary key.""" + coro = self._session.get(model, pk) + if self._is_async: + return coro + return coro + + def add(self, obj: Any) -> None: + """Stage an object for insertion.""" + self._session.add(obj) + + def flush(self) -> Any: + """Flush pending changes to the DB without committing.""" + result = self._session.flush() + if hasattr(result, "__await__"): + return self._maybe_async(result) + return result + + def delete(self, obj: Any) -> Any: + """Mark an object for deletion.""" + result = self._session.delete(obj) + if hasattr(result, "__await__"): + return self._maybe_async(result) + return result + + def refresh(self, obj: Any, attributes: Sequence[str] | None = None) -> Any: + """Re-read object attributes from the DB.""" + kwargs = {} + if attributes: + kwargs["attribute_names"] = list(attributes) + result = self._session.refresh(obj, **kwargs) + if hasattr(result, "__await__"): + return self._maybe_async(result) + return result + + def execute(self, query: Any, *args: Any, **kwargs: Any) -> Any: + """Execute a query object and return the result.""" + result = self._session.execute(query, *args, **kwargs) + if hasattr(result, "__await__"): + return self._maybe_async(result) + return result + + def commit(self) -> Any: + """Persist all pending changes.""" + result = self._session.commit() + if hasattr(result, "__await__"): + return self._maybe_async(result) + return result + + def rollback(self) -> Any: + """Discard all pending changes.""" + result = self._session.rollback() + if hasattr(result, "__await__"): + return self._maybe_async(result) + return result + + def close(self) -> Any: + """Close the underlying session.""" + result = self._session.close() + if hasattr(result, "__await__"): + return self._maybe_async(result) + return result + + +# --------------------------------------------------------------------------- +# #25 — Query Adapter +# --------------------------------------------------------------------------- + + +class SqlAlchemyQueryAdapter: + """Chainable wrapper around SQLAlchemy ``select()`` statements. + + Implements :class:`QueryBackend` via structural subtyping. + """ + + def select(self, model: type) -> Any: + """Start a new query for the given model.""" + from sqlalchemy import select as sa_select + + return sa_select(model) + + def where(self, query: Any, *conditions: Any) -> Any: + """Add WHERE conditions (AND composition).""" + return query.where(*conditions) + + def order_by(self, query: Any, *columns: Any) -> Any: + """Add ORDER BY clauses. Prefix ``-`` for descending.""" + from sqlalchemy import asc, desc + + resolved: list[Any] = [] + for col in columns: + if isinstance(col, str) and col.startswith("-"): + resolved.append(desc(col[1:])) + else: + resolved.append(asc(col) if isinstance(col, str) else col) + return query.order_by(*resolved) + + def limit(self, query: Any, n: int) -> Any: + """Limit the result set to *n* rows.""" + return query.limit(n) + + def offset(self, query: Any, n: int) -> Any: + """Skip the first *n* rows of the result set.""" + return query.offset(n) + + def join(self, query: Any, related: type, on: Any | None = None) -> Any: + """Join a related model onto the query.""" + if on is not None: + return query.join(related, on) + return query.join(related) + + def distinct(self, query: Any) -> Any: + """Add DISTINCT to the query.""" + return query.distinct() + + def count(self, query: Any) -> int: + """Execute the query and return the total row count. + + Wraps the query in a subquery and counts all rows. + """ + from sqlalchemy import func + from sqlalchemy import select as sa_select + + # Extract the selectable from the query + subq = query.subquery() + count_q = sa_select(func.count()).select_from(subq) + # The caller must execute this; return the compiled query + # so the caller can pass it to session.execute() + return count_q + + def options(self, query: Any, *opts: Any) -> Any: + """Add eager-load options (joinedload, selectinload, etc.).""" + return query.options(*opts) + + def ilike(self, column: Any, pattern: str) -> Any: + """Apply case-insensitive LIKE to a column, returning a boolean clause.""" + return column.ilike(pattern) + + def or_(self, *clauses: Any) -> Any: + """Compose multiple boolean clauses with OR.""" + from sqlalchemy import or_ + + return or_(*clauses) + + +# --------------------------------------------------------------------------- +# #30 — Database Backend +# --------------------------------------------------------------------------- + + +class SqlAlchemyDatabaseBackend: + """Wraps ``AdminDatabase``'s engine/table/migration logic to implement + :class:`DatabaseBackend`. + """ + + def __init__( + self, + admin_database: Any | None = None, + database_config: Any | None = None, + ) -> None: + self._admin_database = admin_database + self._database_config = database_config + + def create_connection(self) -> Any: + """Create and return a new SQLAlchemy async engine.""" + if self._admin_database is not None: + self._admin_database._ensure_engine() + return self._admin_database.engine + if self._database_config is not None: + return self._database_config.create_engine() + raise ValueError("No admin_database or database_config provided") + + def create_tables(self, connection: Any, metadata: Any) -> None: + """Issue DDL to create all tables defined in *metadata*. + + For async engines, ``connection`` should be the engine itself; + tables are created via ``run_sync``. + """ + from sqlalchemy.ext.asyncio import AsyncEngine + + if isinstance(connection, AsyncEngine): + import asyncio + + async def _create() -> None: + async with connection.begin() as conn: + await conn.run_sync(metadata.create_all) + + try: + loop = asyncio.get_running_loop() + except RuntimeError: + loop = None + if loop and loop.is_running(): + # We're inside an async context — caller should use run_sync + return _create() + asyncio.run(_create()) + else: + metadata.create_all(bind=connection) + + def auto_migrate(self, connection: Any, metadata: Any) -> None: + """Detect schema drift and add missing columns automatically.""" + from sqlalchemy.ext.asyncio import AsyncEngine + + if isinstance(connection, AsyncEngine): + if self._admin_database is not None: + import asyncio + + async def _migrate() -> None: + async with connection.begin() as conn: + await conn.run_sync(self._admin_database._auto_migrate, metadata) + + try: + loop = asyncio.get_running_loop() + except RuntimeError: + loop = None + if loop and loop.is_running(): + return _migrate() + asyncio.run(_migrate()) + elif self._admin_database is not None: + self._admin_database._auto_migrate_sync(metadata) + + def create_session_factory(self, connection: Any) -> Any: + """Create an ``async_sessionmaker`` bound to *connection*.""" + from fastapi_admin_kit.db import create_session_factory + + return create_session_factory(connection) diff --git a/fastapi_admin_kit/db.py b/fastapi_admin_kit/db.py index 1c68d19..a423280 100644 --- a/fastapi_admin_kit/db.py +++ b/fastapi_admin_kit/db.py @@ -7,6 +7,7 @@ from __future__ import annotations +from collections.abc import Sequence from typing import Any from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker @@ -24,30 +25,35 @@ def create_session_factory( ) -def get_db_session(request: Request) -> AsyncSession: - """Return the per-request ``AsyncSession``. +def _wrap_session(session: Any) -> Any: + """Wrap a raw session in ``SqlAlchemySessionAdapter``.""" + from fastapi_admin_kit.backends.sqlalchemy import SqlAlchemySessionAdapter + + return SqlAlchemySessionAdapter(session) + + +def get_db_session(request: Request) -> Any: + """Return the per-request ``SqlAlchemySessionAdapter`` (implements ``SessionBackend``). The session is created by :class:`SessionMiddleware` and stored on ``scope["state"]["admin_db_session"]`` (accessible via ``request.state.admin_db_session``). Falls back to the legacy ``app.state.admin_db_session`` when the middleware is not active. """ + from fastapi_admin_kit.backends.sqlalchemy import SqlAlchemySessionAdapter + session = getattr(request.state, "admin_db_session", None) if session is not None: - if isinstance(session, AsyncSession): + if isinstance(session, SqlAlchemySessionAdapter): return session - from fastapi_admin_kit.db import SyncSessionWrapper - - return SyncSessionWrapper(session) + return _wrap_session(session) real_app = getattr(request.scope, "app", None) or request.app legacy = getattr(real_app.state, "admin_db_session", None) if legacy is not None: - if isinstance(legacy, AsyncSession): + if isinstance(legacy, SqlAlchemySessionAdapter): return legacy - from fastapi_admin_kit.db import SyncSessionWrapper - - return SyncSessionWrapper(legacy) - return legacy + return _wrap_session(legacy) + return _wrap_session(legacy) # type: ignore[arg-type] class SessionMiddleware: @@ -106,28 +112,50 @@ async def __call__(self, scope: dict, receive: Any, send: Any) -> None: class SyncSessionWrapper: - """Wraps a sync SQLAlchemy Session to provide an async-compatible interface.""" + """Wraps a sync SQLAlchemy Session to provide an async-compatible interface. + + Also implements :class:`SessionBackend` (via ``SqlAlchemySessionAdapter``). + """ def __init__(self, session: Any) -> None: self._session = session + from fastapi_admin_kit.backends.sqlalchemy import SqlAlchemySessionAdapter - async def execute(self, *args: Any, **kwargs: Any) -> Any: - return self._session.execute(*args, **kwargs) + self._adapter = SqlAlchemySessionAdapter(session) - async def commit(self) -> None: - self._session.commit() + @property + def adapter(self) -> Any: + return self._adapter - async def rollback(self) -> None: - self._session.rollback() + def get(self, model: type, pk: Any) -> Any | None: + return self._adapter.get(model, pk) - async def close(self) -> None: - self._session.close() + def add(self, obj: Any) -> None: + self._adapter.add(obj) + + def flush(self) -> None: + self._adapter.flush() + + def delete(self, obj: Any) -> None: + self._adapter.delete(obj) + + def refresh(self, obj: Any, attributes: Sequence[str] | None = None) -> None: + self._adapter.refresh(obj, attributes) + + def commit(self) -> None: + self._adapter.commit() + + def rollback(self) -> None: + self._adapter.rollback() + + def close(self) -> None: + self._adapter.close() + + async def execute(self, *args: Any, **kwargs: Any) -> Any: + return self._session.execute(*args, **kwargs) async def merge(self, *args: Any, **kwargs: Any) -> Any: return self._session.merge(*args, **kwargs) - async def flush(self) -> None: - self._session.flush() - def __getattr__(self, name: str) -> Any: return getattr(self._session, name) diff --git a/fastapi_admin_kit/filters/base.py b/fastapi_admin_kit/filters/base.py index d8ec4c3..2965eca 100644 --- a/fastapi_admin_kit/filters/base.py +++ b/fastapi_admin_kit/filters/base.py @@ -14,45 +14,56 @@ def __init__(self, field_name: str, label: str = "") -> None: self.label = label or field_name.replace("_", " ").title() @abstractmethod - def apply(self, query: Any, value: str) -> Any: - """Apply the filter to a SQLAlchemy select query.""" + def apply(self, query_adapter: Any, query: Any, model: Any, value: str) -> Any: + """Apply the filter to a query via QueryBackend. + + Args: + query_adapter: A QueryBackend adapter instance. + query: The current query statement. + model: The ORM model the query selects. + value: The filter value from request query params. + + Returns: + The updated query with the filter applied. + """ ... def get_choices(self, session: Any) -> list[tuple[str, str]]: - """Return available filter choices as (value, label) pairs.""" + """Return available filter choices as (value, label) pairs. + + Args: + session: A SessionBackend adapter instance. + """ return [] class TextFilter(Filter): """Simple text equality filter.""" - def apply(self, query: Any, value: str) -> Any: - model = query.column_descriptions[0]["entity"] + def apply(self, query_adapter: Any, query: Any, model: Any, value: str) -> Any: if hasattr(model, self.field_name): col = getattr(model, self.field_name) - return query.where(col == value) + return query_adapter.where(query, col == value) return query class BooleanFilter(Filter): """Boolean filter — maps '1' to True, '0' to False.""" - def apply(self, query: Any, value: str) -> Any: - model = query.column_descriptions[0]["entity"] + def apply(self, query_adapter: Any, query: Any, model: Any, value: str) -> Any: if hasattr(model, self.field_name): col = getattr(model, self.field_name) - return query.where(col == (value == "1")) + return query_adapter.where(query, col == (value == "1")) return query class RelationFilter(Filter): """Filter by foreign key relationship.""" - def apply(self, query: Any, value: str) -> Any: - model = query.column_descriptions[0]["entity"] + def apply(self, query_adapter: Any, query: Any, model: Any, value: str) -> Any: if hasattr(model, self.field_name): col = getattr(model, self.field_name) - return query.where(col == value) + return query_adapter.where(query, col == value) return query @@ -68,11 +79,10 @@ def __init__( super().__init__(field_name, label) self._enum_choices = choices or [] - def apply(self, query: Any, value: str) -> Any: - model = query.column_descriptions[0]["entity"] + def apply(self, query_adapter: Any, query: Any, model: Any, value: str) -> Any: if hasattr(model, self.field_name): col = getattr(model, self.field_name) - return query.where(col == value) + return query_adapter.where(query, col == value) return query def get_choices(self, session: Any) -> list[tuple[str, str]]: @@ -85,26 +95,24 @@ def get_choices(self, session: Any) -> list[tuple[str, str]]: class NumericFilter(Filter): """Numeric range filter (gte/lte).""" - def apply(self, query: Any, value: str) -> Any: - model = query.column_descriptions[0]["entity"] + def apply(self, query_adapter: Any, query: Any, model: Any, value: str) -> Any: if not hasattr(model, self.field_name): return query col = getattr(model, self.field_name) if isinstance(value, dict): if value.get("gte"): - query = query.where(col >= value["gte"]) + query = query_adapter.where(query, col >= value["gte"]) if value.get("lte"): - query = query.where(col <= value["lte"]) + query = query_adapter.where(query, col <= value["lte"]) return query class DateRangeFilter(Filter): """Date range filter (from/to).""" - def apply(self, query: Any, value: str) -> Any: + def apply(self, query_adapter: Any, query: Any, model: Any, value: str) -> Any: from datetime import date - model = query.column_descriptions[0]["entity"] if not hasattr(model, self.field_name): return query col = getattr(model, self.field_name) @@ -112,13 +120,13 @@ def apply(self, query: Any, value: str) -> Any: if value.get("from"): try: d = date.fromisoformat(value["from"]) - query = query.where(col >= d) + query = query_adapter.where(query, col >= d) except (ValueError, TypeError): pass if value.get("to"): try: d = date.fromisoformat(value["to"]) - query = query.where(col <= d) + query = query_adapter.where(query, col <= d) except (ValueError, TypeError): pass return query @@ -127,10 +135,9 @@ def apply(self, query: Any, value: str) -> Any: class DatetimeRangeFilter(Filter): """Datetime range filter (from/to).""" - def apply(self, query: Any, value: str) -> Any: + def apply(self, query_adapter: Any, query: Any, model: Any, value: str) -> Any: from datetime import datetime - model = query.column_descriptions[0]["entity"] if not hasattr(model, self.field_name): return query col = getattr(model, self.field_name) @@ -138,13 +145,13 @@ def apply(self, query: Any, value: str) -> Any: if value.get("from"): try: dt = datetime.fromisoformat(value["from"]) - query = query.where(col >= dt) + query = query_adapter.where(query, col >= dt) except (ValueError, TypeError): pass if value.get("to"): try: dt = datetime.fromisoformat(value["to"]) - query = query.where(col <= dt) + query = query_adapter.where(query, col <= dt) except (ValueError, TypeError): pass return query @@ -162,9 +169,8 @@ def __init__( super().__init__(field_name, label) self.search_fields = search_fields or ["name"] - def apply(self, query: Any, value: str) -> Any: - model = query.column_descriptions[0]["entity"] + def apply(self, query_adapter: Any, query: Any, model: Any, value: str) -> Any: if not hasattr(model, self.field_name): return query col = getattr(model, self.field_name) - return query.where(col == value) + return query_adapter.where(query, col == value) diff --git a/fastapi_admin_kit/filters/registry.py b/fastapi_admin_kit/filters/registry.py index 0f66a12..a84496a 100644 --- a/fastapi_admin_kit/filters/registry.py +++ b/fastapi_admin_kit/filters/registry.py @@ -25,10 +25,28 @@ def register(self, model_name: str, filter_obj: Filter) -> None: def get_filters(self, model_name: str) -> dict[str, Filter]: return self._filters.get(model_name, {}).copy() - def auto_generate(self, model: Any, columns: list[Any]) -> dict[str, Filter]: - from sqlalchemy import inspect as sa_inspect + def auto_generate( + self, + model: Any, + columns: list[Any], + introspection: Any | None = None, + ) -> dict[str, Filter]: + """Auto-generate filters for a model's columns. + + Args: + model: The ORM model. + columns: List of ColumnMeta for the model. + introspection: Optional IntrospectionBackend adapter. When None, + falls back to direct SQLAlchemy inspection. + """ + if introspection is not None: + rel_names = introspection.get_relationship_names(model) + else: + from sqlalchemy import inspect as sa_inspect + + mapper = sa_inspect(model) + rel_names = {r.key for r in mapper.relationships} - mapper = sa_inspect(model) filters: dict[str, Filter] = {} for col_meta in columns: @@ -36,27 +54,33 @@ def auto_generate(self, model: Any, columns: list[Any]) -> dict[str, Filter]: if field_name == "id": continue - rel_names = {r.key for r in mapper.relationships} if field_name in rel_names: filters[field_name] = RelationFilter(field_name) continue - for prop in mapper.column_attrs: - if prop.key != field_name: - continue - col = prop.columns[0] if prop.columns else None - if col is None: - break - - type_name = col.type.__class__.__name__ - if type_name == "Boolean": - filters[field_name] = BooleanFilter(field_name) - elif hasattr(col.type, "enums") and col.type.enums: - filters[field_name] = EnumFilter(field_name, choices=list(col.type.enums)) - elif col.foreign_keys: - filters[field_name] = RelationFilter(field_name) - else: - filters[field_name] = TextFilter(field_name) - break + if introspection is not None: + type_name = introspection.get_column_type_name(model, field_name) + col = introspection.get_column_attr(model, field_name) + else: + from sqlalchemy import inspect as sa_inspect + + mapper = sa_inspect(model) + type_name = None + col = None + for prop in mapper.column_attrs: + if prop.key == field_name: + col = prop.columns[0] if prop.columns else None + if col is not None: + type_name = col.type.__class__.__name__ + break + + if type_name == "Boolean": + filters[field_name] = BooleanFilter(field_name) + elif col is not None and hasattr(col.type, "enums") and col.type.enums: + filters[field_name] = EnumFilter(field_name, choices=list(col.type.enums)) + elif col is not None and col.foreign_keys: + filters[field_name] = RelationFilter(field_name) + else: + filters[field_name] = TextFilter(field_name) return filters diff --git a/fastapi_admin_kit/inspection/__init__.py b/fastapi_admin_kit/inspection/__init__.py index 646518e..c8ea147 100644 --- a/fastapi_admin_kit/inspection/__init__.py +++ b/fastapi_admin_kit/inspection/__init__.py @@ -1,54 +1,28 @@ -"""Model inspection — SQLAlchemy model → ColumnMeta / RelationMeta.""" +"""Model inspection — SQLAlchemy model → ColumnMeta / RelationMeta. + +Re-exports backward-compatible module-level functions from +:class:`~fastapi_admin_kit.backends.sqlalchemy.SqlAlchemyIntrospectionAdapter`. +""" from __future__ import annotations import re from typing import Any -from sqlalchemy import inspect - +from fastapi_admin_kit.backends.sqlalchemy import SqlAlchemyIntrospectionAdapter from fastapi_admin_kit.types import ColumnMeta, RelationMeta +_inspector = SqlAlchemyIntrospectionAdapter() + def inspect_model(model: type) -> tuple[list[ColumnMeta], list[RelationMeta]]: """Inspect a SQLAlchemy model and return column + relationship metadata.""" - mapper = inspect(model) - columns: list[ColumnMeta] = [] - relationships: list[RelationMeta] = [] - - for col in mapper.columns: - columns.append( - ColumnMeta( - name=col.key, - type=col.type, - nullable=col.nullable, - primary_key=col.primary_key, - foreign_keys=list(col.foreign_keys), - default=col.default, - server_default=col.server_default, - index=col.index, - unique=col.unique, - ) - ) - - for rel in mapper.relationships: - relationships.append( - RelationMeta( - name=rel.key, - direction=rel.direction.name, - target_model=rel.mapper.class_, - uselist=rel.uselist, - back_populates=rel.back_populates, - secondary=rel.secondary, - ) - ) - - return columns, relationships + return _inspector.inspect_model(model) def is_abstract(model: type) -> bool: """Check if a model is abstract and should be skipped during auto-discovery.""" - return getattr(model, "__abstract__", False) + return _inspector.is_abstract(model) def get_pk_field(model: type) -> str | None: @@ -58,13 +32,7 @@ def get_pk_field(model: type) -> str | None: or a tuple of names for composite PKs. Returns None if no primary key is found. """ - mapper = inspect(model) - pk_cols = mapper.primary_key - if not pk_cols: - return None - if len(pk_cols) == 1: - return pk_cols[0].key - return tuple(col.key for col in pk_cols) + return _inspector.get_pk_field(model) def cast_pk_value(model: type, value: Any) -> Any: @@ -74,25 +42,7 @@ def cast_pk_value(model: type, value: Any) -> Any: accordingly. Supports Integer, BigInteger, String, and UUID types. Returns the original value if type cannot be determined. """ - if value is None: - return None - mapper = inspect(model) - pk_cols = mapper.primary_key - if not pk_cols or len(pk_cols) != 1: - return value - pk_col = pk_cols[0] - from sqlalchemy import BigInteger, Integer - from sqlalchemy.dialects.postgresql import UUID as PG_UUID - from sqlalchemy.types import Uuid - - col_type = type(pk_col.type) - if col_type in (Integer, BigInteger): - return int(value) - if col_type in (PG_UUID, Uuid): - from uuid import UUID - - return UUID(str(value)) - return value + return _inspector.cast_pk_value(model, value) def cast_value(col_meta: Any, value: Any) -> Any: diff --git a/fastapi_admin_kit/inspection/registry.py b/fastapi_admin_kit/inspection/registry.py index a3f4558..3d23ee3 100644 --- a/fastapi_admin_kit/inspection/registry.py +++ b/fastapi_admin_kit/inspection/registry.py @@ -3,23 +3,22 @@ from __future__ import annotations import re -from typing import TYPE_CHECKING, Any - -from sqlalchemy import inspect +from typing import Any +from fastapi_admin_kit.backends.sqlalchemy import SqlAlchemyIntrospectionAdapter from fastapi_admin_kit.types import ColumnMeta, RelationMeta -if TYPE_CHECKING: - pass - class ModelInspector: """Inspects SQLAlchemy models and extracts column/relationship metadata. - This class centralizes all model inspection logic, making it testable - and separable from the registry's storage concerns. + Delegates core inspection to :class:`SqlAlchemyIntrospectionAdapter` + and adds validation / metadata-extraction helpers on top. """ + def __init__(self) -> None: + self._adapter = SqlAlchemyIntrospectionAdapter() + def inspect_model(self, model: type) -> tuple[list[ColumnMeta], list[RelationMeta]]: """Inspect a SQLAlchemy or SQLModel model and return column + relationship metadata. @@ -29,113 +28,7 @@ def inspect_model(self, model: type) -> tuple[list[ColumnMeta], list[RelationMet Returns: A tuple of (columns, relationships) metadata. """ - mapper = inspect(model) - columns: list[ColumnMeta] = [] - relationships: list[RelationMeta] = [] - - # Check if this is a SQLModel - is_sqlmodel = self._is_sqlmodel(model) - - for col in mapper.columns: - # For SQLModel, we may need to extract type info from Pydantic fields - col_type = col.type - if is_sqlmodel: - col_type = self._resolve_sqlmodel_type(model, col.key, col.type) - - columns.append( - ColumnMeta( - name=col.key, - type=col_type, - nullable=col.nullable, - primary_key=col.primary_key, - foreign_keys=list(col.foreign_keys), - default=col.default, - server_default=col.server_default, - index=col.index, - unique=col.unique, - ) - ) - - for rel in mapper.relationships: - try: - relationships.append( - RelationMeta( - name=rel.key, - direction=rel.direction.name, - target_model=rel.mapper.class_, - uselist=rel.uselist, - back_populates=rel.back_populates, - secondary=rel.secondary, - ) - ) - except Exception: - # Skip relationships that fail to configure (e.g. SQLModel - # relationships with incomplete FK resolution) - pass - - return columns, relationships - - def _is_sqlmodel(self, model: type) -> bool: - """Check if a model is a SQLModel instance.""" - try: - from sqlmodel import SQLModel - - return isinstance(model, type) and issubclass(model, SQLModel) - except ImportError: - return False - - def _resolve_sqlmodel_type(self, model: type, field_name: str, default_type: Any) -> Any: - """Resolve the column type for a SQLModel field. - - SQLModel may expose Python types (int, str, etc.) instead of SQLAlchemy types. - This method maps them to equivalent SQLAlchemy types. - """ - try: - from sqlmodel import SQLModel - - if not (isinstance(model, type) and issubclass(model, SQLModel)): - return default_type - - # Get SQLModel field info - sqlmodel_fields = getattr(model, "__sqlmodel_fields__", {}) - if field_name not in sqlmodel_fields: - return default_type - - field_info = sqlmodel_fields[field_name] - annotation = getattr(field_info, "annotation", None) - - if annotation is None: - return default_type - - # Map Python types to SQLAlchemy types - type_map = { - int: self._get_sa_type("Integer"), - str: self._get_sa_type("String"), - float: self._get_sa_type("Float"), - bool: self._get_sa_type("Boolean"), - } - - # Handle Optional types - origin = getattr(annotation, "__origin__", None) - if origin is not None: - args = getattr(annotation, "__args__", ()) - if args: - inner = args[0] - if inner in type_map: - return type_map[inner] - - if annotation in type_map: - return type_map[annotation] - - return default_type - except Exception: - return default_type - - def _get_sa_type(self, type_name: str) -> Any: - """Get a SQLAlchemy type by name.""" - import sqlalchemy as sa - - return getattr(sa, type_name, None) + return self._adapter.inspect_model(model) def validate_model(self, model: type) -> None: """Validate that a model is suitable for admin registration. @@ -184,7 +77,7 @@ def is_abstract(self, model: type) -> bool: Returns: True if the model is abstract, False otherwise. """ - return getattr(model, "__abstract__", False) + return self._adapter.is_abstract(model) def get_pk_field(self, model: type) -> str | tuple[str, ...] | None: """Get the primary key field name for a model. @@ -197,13 +90,7 @@ def get_pk_field(self, model: type) -> str | tuple[str, ...] | None: a tuple of names for composite PKs, or None if no primary key is found. """ - mapper = inspect(model) - pk_cols = mapper.primary_key - if not pk_cols: - return None - if len(pk_cols) == 1: - return pk_cols[0].key - return tuple(col.key for col in pk_cols) + return self._adapter.get_pk_field(model) def auto_label(self, name: str) -> str: """Auto-generate a human-readable label from a field name. diff --git a/fastapi_admin_kit/router.py b/fastapi_admin_kit/router.py index acacc8b..63071a0 100644 --- a/fastapi_admin_kit/router.py +++ b/fastapi_admin_kit/router.py @@ -482,7 +482,7 @@ async def autocomplete( from fastapi_admin_kit.search_utils import apply_search_filter - query = apply_search_filter(select(model), model, search_fields, q).limit(20) + query = apply_search_filter(request, select(model), model, search_fields, q).limit(20) result = await session.execute(query) for obj in result.scalars(): label = str( diff --git a/fastapi_admin_kit/search_utils.py b/fastapi_admin_kit/search_utils.py index 6a1f902..a9766e2 100644 --- a/fastapi_admin_kit/search_utils.py +++ b/fastapi_admin_kit/search_utils.py @@ -11,15 +11,20 @@ def apply_search_filter( + request: Any, query: Any, model: Any, search_fields: list[str] | None, q: str, ) -> Any: - """Apply case-insensitive ``ilike`` search clauses to a SQLAlchemy query. + """Apply case-insensitive ``ilike`` search clauses to a query. + + Uses ``QueryBackend`` and ``IntrospectionBackend`` from ``app.state`` + when available, falling back to direct SQLAlchemy imports. Args: - query: A SQLAlchemy ``select`` statement (or query object). + request: The FastAPI Request (used to access ``app.state`` backends). + query: A query statement (SQLAlchemy ``select`` or backend query type). model: The root ORM model the query selects. search_fields: Field names to search. Plain names match direct columns; names containing ``__`` (e.g. ``roles__name``) match an attribute on @@ -34,17 +39,32 @@ def apply_search_filter( if not q or not search_fields: return query - from sqlalchemy import inspect as sa_inspect - from sqlalchemy import or_ + query_adapter = getattr(request.app.state, "admin_query_adapter", None) + introspection = getattr(request.app.state, "admin_introspection_adapter", None) + + if introspection is not None: + rel_names = introspection.get_relationship_names(model) + else: + from sqlalchemy import inspect as sa_inspect + + mapper = sa_inspect(model) + rel_names = {r.key for r in mapper.relationships} - mapper = sa_inspect(model) clauses: list[Any] = [] joined_rels: set[str] = set() for sf in search_fields: if "__" in sf: rel_name, attr = sf.split("__", 1) - rel = mapper.relationships.get(rel_name) + if rel_name not in rel_names: + continue + if introspection is not None: + rel = introspection.get_relationship(model, rel_name) + else: + from sqlalchemy import inspect as sa_inspect + + mapper = sa_inspect(model) + rel = mapper.relationships.get(rel_name) if rel is None: continue target = rel.mapper.class_ @@ -54,18 +74,35 @@ def apply_search_filter( if not hasattr(col, "ilike"): continue if rel_name not in joined_rels: - query = query.join(getattr(model, rel_name)) + if query_adapter is not None: + query = query_adapter.join(query, getattr(model, rel_name)) + else: + query = query.join(getattr(model, rel_name)) joined_rels.add(rel_name) - clauses.append(col.ilike(f"%{q}%")) + if query_adapter is not None: + clauses.append(query_adapter.ilike(col, f"%{q}%")) + else: + clauses.append(col.ilike(f"%{q}%")) else: if hasattr(model, sf): col = getattr(model, sf) if hasattr(col, "ilike"): - clauses.append(col.ilike(f"%{q}%")) + if query_adapter is not None: + clauses.append(query_adapter.ilike(col, f"%{q}%")) + else: + clauses.append(col.ilike(f"%{q}%")) if clauses: - query = query.where(or_(*clauses)) + if query_adapter is not None: + query = query_adapter.where(query, query_adapter.or_(*clauses)) + else: + from sqlalchemy import or_ + + query = query.where(or_(*clauses)) if joined_rels: - query = query.distinct() + if query_adapter is not None: + query = query_adapter.distinct(query) + else: + query = query.distinct() return query diff --git a/fastapi_admin_kit/views/class_views.py b/fastapi_admin_kit/views/class_views.py index 2955683..a74d307 100644 --- a/fastapi_admin_kit/views/class_views.py +++ b/fastapi_admin_kit/views/class_views.py @@ -219,7 +219,7 @@ async def _build_filter_fields(self, request: Request) -> dict[str, dict[str, An filter_fields: dict[str, dict[str, Any]] = {} for filter_field in self.admin.list_filter: filter_fields[filter_field] = await self.query_provider._get_filter_choices( - model, filter_field, session + request, model, filter_field, session ) return filter_fields @@ -1354,7 +1354,7 @@ async def _search(self, request: Request, q: str, limit: int = 20, exclude_id: s "name", "title", ] - base = apply_search_filter(select(model), model, search_fields, q) + base = apply_search_filter(request, select(model), model, search_fields, q) if exclude_id: pk_col = getattr(model, self.registered.pk_field, None) diff --git a/fastapi_admin_kit/views/context.py b/fastapi_admin_kit/views/context.py index 86f1a9b..c2e4302 100644 --- a/fastapi_admin_kit/views/context.py +++ b/fastapi_admin_kit/views/context.py @@ -351,7 +351,7 @@ async def build_list_context( base = base.where(and_(*filter_clauses)) if q and registered.admin.search_fields: - base = apply_search_filter(base, model, registered.admin.search_fields, q) + base = apply_search_filter(request, base, model, registered.admin.search_fields, q) query_ordering = request.query_params.get("ordering", "") if query_ordering: diff --git a/fastapi_admin_kit/views/renderers.py b/fastapi_admin_kit/views/renderers.py index 230c3a0..c8cdca3 100644 --- a/fastapi_admin_kit/views/renderers.py +++ b/fastapi_admin_kit/views/renderers.py @@ -257,71 +257,132 @@ async def parse( class DefaultQueryProvider: - """SRP: Build and execute SQLAlchemy queries with filtering, search, pagination.""" + """SRP: Build and execute queries with filtering, search, pagination.""" def __init__(self, registered: RegisteredModel): self.registered = registered - def _get_eager_loads(self, model: Any, list_display: list[str]) -> list: + def _get_query_adapter(self, request: Request) -> Any: + """Get the QueryBackend from app.state.""" + return getattr(request.app.state, "admin_query_adapter", None) + + def _get_introspection(self, request: Request) -> Any: + """Get the IntrospectionBackend from app.state.""" + return getattr(request.app.state, "admin_introspection_adapter", None) + + def _get_eager_loads(self, request: Request, model: Any, list_display: list[str]) -> list: """Build eager load options for relationship columns.""" - from sqlalchemy import inspect as sa_inspect from sqlalchemy.orm import joinedload - mapper = sa_inspect(model) - rel_names = {r.key for r in mapper.relationships} + introspection = self._get_introspection(request) + if introspection is not None: + rel_names = introspection.get_relationship_names(model) + else: + from sqlalchemy import inspect as sa_inspect + + mapper = sa_inspect(model) + rel_names = {r.key for r in mapper.relationships} options = [] for col_name in list_display: if col_name in rel_names: options.append(joinedload(getattr(model, col_name))) return options - def _get_field_type(self, model: Any, field_name: str) -> str: + def _get_field_type(self, request: Request, model: Any, field_name: str) -> str: """Detect the abstract field type for a model field.""" - from sqlalchemy import inspect as sa_inspect + introspection = self._get_introspection(request) + if introspection is not None: + rel_names = introspection.get_relationship_names(model) + else: + from sqlalchemy import inspect as sa_inspect - mapper = sa_inspect(model) - rel_names = {r.key for r in mapper.relationships} + mapper = sa_inspect(model) + rel_names = {r.key for r in mapper.relationships} if field_name in rel_names: return "relation" - for prop in mapper.column_attrs: - if prop.key == field_name: - col = prop.columns[0] if prop.columns else None - if col is None: + if introspection is not None: + type_name = introspection.get_column_type_name(model, field_name) + else: + from sqlalchemy import inspect as sa_inspect + + mapper = sa_inspect(model) + type_name = None + for prop in mapper.column_attrs: + if prop.key == field_name: + col = prop.columns[0] if prop.columns else None + if col is not None: + type_name = col.type.__class__.__name__ + break + + if type_name == "Boolean": + return "boolean" + if type_name == "DateTime": + return "datetime" + if type_name == "Date": + return "date" + if type_name == "Time": + return "time" + + if introspection is not None: + col = introspection.get_column_attr(model, field_name) + if col is not None and hasattr(col.type, "enums") and col.type.enums: + return "enum" + if col is not None and col.foreign_keys: + return "relation" + else: + from sqlalchemy import inspect as sa_inspect + + mapper = sa_inspect(model) + for prop in mapper.column_attrs: + if prop.key == field_name: + col = prop.columns[0] if prop.columns else None + if col is not None: + if hasattr(col.type, "enums") and col.type.enums: + return "enum" + if col.foreign_keys: + return "relation" break - type_name = col.type.__class__.__name__ - if type_name == "Boolean": - return "boolean" - if type_name == "DateTime": - return "datetime" - if type_name == "Date": - return "date" - if type_name == "Time": - return "time" - if hasattr(col.type, "enums") and col.type.enums: - return "enum" - if col.foreign_keys: - return "relation" - return "text" + return "text" async def _get_filter_choices( - self, model: Any, field_name: str, session: Any = None + self, request: Request, model: Any, field_name: str, session: Any = None ) -> dict[str, Any]: """Get filter field type and available choices for a field.""" - from sqlalchemy import inspect as sa_inspect - from sqlalchemy import select - - mapper = sa_inspect(model) - field_type = self._get_field_type(model, field_name) + introspection = self._get_introspection(request) + field_type = self._get_field_type(request, model, field_name) if field_type == "relation": - rel_map = {r.key: r for r in mapper.relationships} + if introspection is not None: + rel = introspection.get_relationship(model, field_name) + else: + from sqlalchemy import inspect as sa_inspect + + mapper = sa_inspect(model) + rel = mapper.relationships.get(field_name) + target_model = None - if field_name in rel_map: - target_model = rel_map[field_name].mapper.class_ + if rel is not None: + target_model = rel.mapper.class_ + elif introspection is not None: + mapper_rel_names = introspection.get_relationship_names(model) + for rname in mapper_rel_names: + r = introspection.get_relationship(model, rname) + if r is not None and r.direction.name == "MANYTOONE": + col = introspection.get_column_attr(model, field_name) + if col is not None: + for fk in col.foreign_keys: + if fk.column.table == r.mapper.persist_selectable: + target_model = r.mapper.class_ + break + if target_model is not None: + break else: + from sqlalchemy import inspect as sa_inspect + + mapper = sa_inspect(model) for rel in mapper.relationships: if rel.direction.name == "MANYTOONE": for prop in mapper.column_attrs: @@ -341,11 +402,35 @@ async def _get_filter_choices( order_col = getattr(target_model, "name", None) or getattr( target_model, "title", None ) - if order_col is not None: - q = select(target_model).order_by(order_col).limit(100) + if introspection is not None: + query_adapter = self._get_query_adapter(request) + else: + query_adapter = None + + if query_adapter is not None: + q = query_adapter.select(target_model) + if order_col is not None: + q = query_adapter.order_by(q, order_col) + else: + if introspection is not None: + pk_cols = introspection.get_pk_columns(target_model) + q = query_adapter.order_by(q, pk_cols[0]) + else: + from sqlalchemy import inspect as sa_inspect + + pk = sa_inspect(target_model).primary_key[0] + q = query_adapter.order_by(q, pk) + q = query_adapter.limit(q, 100) else: - pk = sa_inspect(target_model).primary_key[0] - q = select(target_model).order_by(pk).limit(100) + from sqlalchemy import inspect as sa_inspect + from sqlalchemy import select + + if order_col is not None: + q = select(target_model).order_by(order_col).limit(100) + else: + pk = sa_inspect(target_model).primary_key[0] + q = select(target_model).order_by(pk).limit(100) + result = await session.execute(q) for obj in result.scalars(): label = str( @@ -365,37 +450,66 @@ async def _get_filter_choices( } if field_type == "enum": - for prop in mapper.column_attrs: - if prop.key == field_name: - col = prop.columns[0] if prop.columns else None - if col is not None and hasattr(col.type, "enums"): - choices = [("", "All")] - for val in col.type.enums: - choices.append((val, val.replace("_", " ").title())) - return {"field_type": field_type, "choices": choices} + if introspection is not None: + col = introspection.get_column_attr(model, field_name) + if col is not None and hasattr(col.type, "enums"): + choices = [("", "All")] + for val in col.type.enums: + choices.append((val, val.replace("_", " ").title())) + return {"field_type": field_type, "choices": choices} + else: + from sqlalchemy import inspect as sa_inspect + + mapper = sa_inspect(model) + for prop in mapper.column_attrs: + if prop.key == field_name: + col = prop.columns[0] if prop.columns else None + if col is not None and hasattr(col.type, "enums"): + choices = [("", "All")] + for val in col.type.enums: + choices.append((val, val.replace("_", " ").title())) + return {"field_type": field_type, "choices": choices} if field_type in ("date", "datetime", "time"): return {"field_type": field_type, "choices": [("", "All")]} choices = [("", "All")] - for prop in mapper.column_attrs: - if prop.key == field_name: - col = prop.columns[0] if prop.columns else None - if col is not None and session is not None: - try: - q = ( - select(col) - .where(col.isnot(None)) - .group_by(col) - .order_by(col) - .limit(100) - ) - result = session.execute(q) - for (val,) in result: - label = str(val).replace("_", " ").title() - choices.append((str(val), label)) - except Exception: - pass + if introspection is not None: + col = introspection.get_column_attr(model, field_name) + if col is not None and session is not None: + try: + from sqlalchemy import select + + q = select(col).where(col.isnot(None)).group_by(col).order_by(col).limit(100) + result = session.execute(q) + for (val,) in result: + label = str(val).replace("_", " ").title() + choices.append((str(val), label)) + except Exception: + pass + else: + from sqlalchemy import inspect as sa_inspect + from sqlalchemy import select + + mapper = sa_inspect(model) + for prop in mapper.column_attrs: + if prop.key == field_name: + col = prop.columns[0] if prop.columns else None + if col is not None and session is not None: + try: + q = ( + select(col) + .where(col.isnot(None)) + .group_by(col) + .order_by(col) + .limit(100) + ) + result = session.execute(q) + for (val,) in result: + label = str(val).replace("_", " ").title() + choices.append((str(val), label)) + except Exception: + pass return {"field_type": "text", "choices": choices} async def get_list( @@ -405,22 +519,31 @@ async def get_list( Returns (items, total, page, per_page). """ - from sqlalchemy import and_, asc, desc, select - from fastapi_admin_kit.search_utils import apply_search_filter session = get_db_session(request) registered = self.registered model = registered.model - base = select(model) + + query_adapter = self._get_query_adapter(request) + if query_adapter is not None: + base = query_adapter.select(model) + else: + from sqlalchemy import select + + base = select(model) list_display = registered.admin.list_display or [ c.name for c in registered.columns if c.name != "id" ] - eager_loads = self._get_eager_loads(model, list_display) - for opt in eager_loads: - base = base.options(opt) + eager_loads = self._get_eager_loads(request, model, list_display) + if query_adapter is not None: + for opt in eager_loads: + base = query_adapter.options(base, opt) + else: + for opt in eager_loads: + base = base.options(opt) if registered.admin.list_filter: filter_clauses = [] @@ -428,7 +551,7 @@ async def get_list( param_key = f"filter_{filter_field}" filter_value = request.query_params.get(param_key, "") if filter_value and hasattr(model, filter_field): - field_type = self._get_field_type(model, filter_field) + field_type = self._get_field_type(request, model, filter_field) col = getattr(model, filter_field) if field_type == "boolean": @@ -490,7 +613,7 @@ async def get_list( if (from_val or to_val) and hasattr(model, filter_field): col = getattr(model, filter_field) - field_type = self._get_field_type(model, filter_field) + field_type = self._get_field_type(request, model, filter_field) if field_type == "date" and from_val: try: from datetime import date as _date @@ -525,10 +648,15 @@ async def get_list( pass if filter_clauses: - base = base.where(and_(*filter_clauses)) + if query_adapter is not None: + base = query_adapter.where(base, *filter_clauses) + else: + from sqlalchemy import and_ + + base = base.where(and_(*filter_clauses)) if q and registered.admin.search_fields: - base = apply_search_filter(base, model, registered.admin.search_fields, q) + base = apply_search_filter(request, base, model, registered.admin.search_fields, q) query_ordering = request.query_params.get("ordering", "") if query_ordering: @@ -539,7 +667,15 @@ async def get_list( col_name = order[0].lstrip("-") col = getattr(model, col_name, None) if hasattr(model, col_name) else None if col is not None: - base = base.order_by(desc(col) if order[0].startswith("-") else asc(col)) + if query_adapter is not None: + if order[0].startswith("-"): + base = query_adapter.order_by(base, f"-{col_name}") + else: + base = query_adapter.order_by(base, col_name) + else: + from sqlalchemy import asc, desc + + base = base.order_by(desc(col) if order[0].startswith("-") else asc(col)) per_page = registered.admin.per_page @@ -570,21 +706,55 @@ async def get_list( async def get_object(self, request: Request, id: Any) -> Any | None: """Return a single object by primary key, eagerly loading M2M relationships.""" - from sqlalchemy import inspect as sa_inspect - from sqlalchemy.orm import selectinload + from fastapi_admin_kit.inspection import cast_pk_value session = get_db_session(request) - mapper = sa_inspect(self.registered.model) - options = [] - for rel in mapper.relationships: - if rel.direction.name == "MANYTOMANY": - options.append(selectinload(getattr(self.registered.model, rel.key))) - from fastapi_admin_kit.inspection import cast_pk_value + introspection = self._get_introspection(request) + query_adapter = self._get_query_adapter(request) + + if introspection is not None: + mapper_rel_names = introspection.get_relationship_names(self.registered.model) + else: + from sqlalchemy import inspect as sa_inspect + + mapper = sa_inspect(self.registered.model) + mapper_rel_names = {r.key for r in mapper.relationships} + + m2m_rel_names = set() + for rel_name in mapper_rel_names: + if introspection is not None: + rel = introspection.get_relationship(self.registered.model, rel_name) + else: + from sqlalchemy import inspect as sa_inspect + + mapper = sa_inspect(self.registered.model) + rel = mapper.relationships.get(rel_name) + if rel is not None and rel.direction.name == "MANYTOMANY": + m2m_rel_names.add(rel_name) int_id = cast_pk_value(self.registered.model, id) - if options: + + if m2m_rel_names and query_adapter is not None: + from sqlalchemy.orm import selectinload + + options = [selectinload(getattr(self.registered.model, rn)) for rn in m2m_rel_names] + stmt = query_adapter.select(self.registered.model) + for opt in options: + stmt = query_adapter.options(stmt, opt) + stmt = query_adapter.where( + stmt, + getattr(self.registered.model, self.registered.pk_field) == int_id, + ) + result = await session.execute(stmt) + return result.scalar_one_or_none() + elif m2m_rel_names: + from sqlalchemy import inspect as sa_inspect from sqlalchemy import select + from sqlalchemy.orm import selectinload + mapper = sa_inspect(self.registered.model) + m2m_rels = [r for r in mapper.relationships if r.direction.name == "MANYTOMANY"] + options = [selectinload(getattr(self.registered.model, r.key)) for r in m2m_rels] stmt = ( select(self.registered.model) .options(*options)