Source code for sqlspec.extensions.starlette.middleware

from contextlib import asynccontextmanager
from typing import TYPE_CHECKING, Any

from starlette.middleware.base import BaseHTTPMiddleware

from sqlspec.core import CorrelationExtractor
from sqlspec.core.sqlcommenter import SQLCommenterContext
from sqlspec.extensions.starlette._utils import get_state_value, pop_state_value, set_state_value
from sqlspec.utils.correlation import CorrelationContext
from sqlspec.utils.sync_tools import ensure_async_, with_ensure_async_
from sqlspec.utils.type_guards import has_name

if TYPE_CHECKING:
    from collections.abc import AsyncIterator

    from starlette.requests import Request
    from starlette.responses import Response

    from sqlspec.extensions.starlette._state import SQLSpecConfigState

__all__ = ("CorrelationMiddleware", "SQLCommenterMiddleware", "SQLSpecAutocommitMiddleware", "SQLSpecManualMiddleware")


class SQLSpecManualMiddleware(BaseHTTPMiddleware):
    """Middleware for manual transaction mode.

    Acquires connection from pool, stores in request.state, releases after request.
    No automatic commit or rollback - user code must handle transactions.
    """

    def __init__(self, app: Any, config_state: "SQLSpecConfigState") -> None:
        """Initialize middleware.

        Args:
            app: Starlette application instance.
            config_state: Configuration state for this database.
        """
        super().__init__(app)
        self.config_state = config_state
async def dispatch(self, request: "Request", call_next: Any) -> Any: """Process request with manual transaction mode. Args: request: Incoming HTTP request. call_next: Next middleware or route handler. Returns: HTTP response. """ async with self._connection_cm(request): return await call_next(request)