Source code for sqlspec.extensions.fastapi.extension

from typing import TYPE_CHECKING, Any, overload

from fastapi import Request

from sqlspec.extensions.fastapi.providers import DEPENDENCY_DEFAULTS
from sqlspec.extensions.fastapi.providers import provide_filters as _provide_filters
from sqlspec.extensions.starlette.extension import SQLSpecPlugin as _StarlettePlugin

if TYPE_CHECKING:
    from collections.abc import Callable

    from sqlspec.config import AsyncDatabaseConfig, SyncDatabaseConfig
    from sqlspec.core import FilterTypes
    from sqlspec.driver import AsyncDriverAdapterBase, SyncDriverAdapterBase
    from sqlspec.extensions.fastapi.providers import DependencyDefaults, FilterConfig

    # Type aliases for static analysis - IDEs see the real types
    _AsyncSession = AsyncDriverAdapterBase
    _SyncSession = SyncDriverAdapterBase
    _Session = AsyncDriverAdapterBase | SyncDriverAdapterBase
else:
    # Runtime fallback - FastAPI sees Any (avoids NameError)
    _AsyncSession = Any
    _SyncSession = Any
    _Session = Any

__all__ = ("SQLSpecPlugin",)


class SQLSpecPlugin(_StarlettePlugin):
    """SQLSpec integration for FastAPI applications.

    Extends Starlette integration with dependency injection helpers for FastAPI's
    Depends() system.
    """

    def _extract_extension_settings(self, config: Any) -> "dict[str, Any]":
        """Extract FastAPI settings from config.extension_config.

        Args:
            config: Database configuration instance.

        Returns:
            Dictionary of FastAPI-specific settings.
        """
        framework_config = config.extension_config.get("fastapi", {})
        pool_key = framework_config.get("pool_key", "db_pool")
        if not config.supports_connection_pooling and pool_key == "db_pool":
            pool_key = f"_db_pool_{id(config)}"
        correlation_headers = framework_config.get("correlation_headers")
        return {
            "connection_key": framework_config.get("connection_key", "db_connection"),
            "pool_key": pool_key,
            "session_key": framework_config.get("session_key", "db_session"),
            "commit_mode": framework_config.get("commit_mode", "manual"),
            "extra_commit_statuses": framework_config.get("extra_commit_statuses"),
            "extra_rollback_statuses": framework_config.get("extra_rollback_statuses"),
            "disable_di": framework_config.get("disable_di", False),
            "enable_correlation_middleware": framework_config.get("enable_correlation_middleware", False),
            "correlation_header": framework_config.get("correlation_header", "x-request-id"),
            "correlation_headers": tuple(correlation_headers) if correlation_headers is not None else None,
            "auto_trace_headers": framework_config.get("auto_trace_headers", True),
            "enable_sqlcommenter_middleware": framework_config.get("enable_sqlcommenter_middleware", True),
            "sqlcommenter_framework": framework_config.get("sqlcommenter_framework", "fastapi"),
        }

    @overload
    def provide_session(
        self, key: None = None
    ) -> "Callable[[Request], AsyncDriverAdapterBase | SyncDriverAdapterBase]": ...

    @overload
    def provide_session(self, key: str) -> "Callable[[Request], AsyncDriverAdapterBase | SyncDriverAdapterBase]": ...

    @overload
    def provide_session(self, key: "type[AsyncDatabaseConfig]") -> "Callable[[Request], AsyncDriverAdapterBase]": ...

    @overload
    def provide_session(self, key: "type[SyncDatabaseConfig]") -> "Callable[[Request], SyncDriverAdapterBase]": ...

    @overload
    def provide_session(self, key: "AsyncDatabaseConfig") -> "Callable[[Request], AsyncDriverAdapterBase]": ...

    @overload
    def provide_session(self, key: "SyncDatabaseConfig") -> "Callable[[Request], SyncDriverAdapterBase]": ...

    def provide_session(
        self,
        key: "str | type[AsyncDatabaseConfig | SyncDatabaseConfig] | AsyncDatabaseConfig | SyncDatabaseConfig | None" = None,
    ) -> "Callable[[Request], AsyncDriverAdapterBase | SyncDriverAdapterBase]":
        """Create dependency factory for session injection.

        Returns a callable that can be used with FastAPI's Depends() to inject
        a database session into route handlers.

        Args:
            key: Optional session key (str), config type for type narrowing, or None.

        Returns:
            Dependency callable for FastAPI Depends().
        """
        # Extract string key if provided, ignore config types/instances (used only for type narrowing)
        session_key = key if isinstance(key, str) or key is None else None

        def dependency(request: Request) -> _Session:
            return self.get_session(request, session_key)  # type: ignore[no-any-return]

        return dependency
def provide_async_session(self, key: "str | None" = None) -> "Callable[[Request], AsyncDriverAdapterBase]": """Create dependency factory for async session injection. Type-narrowed version of provide_session() that returns AsyncDriverAdapterBase. Useful when using string keys and you know the config is async. Args: key: Optional session key for multi-database configurations. Returns: Dependency callable that returns AsyncDriverAdapterBase. """ def dependency(request: Request) -> _AsyncSession: return self.get_session(request, key) # type: ignore[no-any-return] return dependency