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