# SPDX-FileCopyrightText: 2023 Hynek Schlawack <hs@ox.cx>
#
# SPDX-License-Identifier: MIT
from __future__ import annotations
import contextlib
import inspect
from collections.abc import AsyncGenerator, Callable
from typing import TYPE_CHECKING, Any, cast, overload
import attrs
from starlette.applications import Starlette
from starlette.requests import Request
from starlette.types import ASGIApp, Receive, Scope, Send
import svcs
from svcs._core import (
_KEY_CONTAINER,
_KEY_REGISTRY,
T1,
T2,
T3,
T4,
T5,
T6,
T7,
T8,
T9,
T10,
TypeForm,
_ServiceType,
)
if TYPE_CHECKING:
from starlette.testclient import TestClient
else:
try:
from starlette.testclient import TestClient
except (ImportError, RuntimeError): # pragma: no cover
TestClient = Any
[docs]
def svcs_from(request: Request) -> svcs.Container:
"""
Get the current container from *request*.
"""
return getattr(request.state, _KEY_CONTAINER) # type: ignore[no-any-return]
[docs]
@attrs.define
class lifespan: # noqa: N801
"""
Make a Starlette lifespan *svcs*-aware.
Makes sure that the registry is available to the decorated lifespan
function as a second parameter and that the registry is closed when the
application exists.
Async generators are automatically wrapped into an async context manager.
Args:
lifespan: The lifespan function to make *svcs*-aware.
"""
_lifespan: (
Callable[
[Starlette, svcs.Registry],
contextlib.AbstractAsyncContextManager[dict[str, object]],
]
| Callable[
[Starlette, svcs.Registry],
contextlib.AbstractAsyncContextManager[None],
]
| Callable[
[Starlette, svcs.Registry], AsyncGenerator[dict[str, object], None]
]
| Callable[[Starlette, svcs.Registry], AsyncGenerator[None, None]]
)
_state: dict[str, object] = attrs.field(factory=dict)
registry: svcs.Registry = attrs.field(factory=svcs.Registry)
@contextlib.asynccontextmanager
async def __call__(
self, app: Starlette
) -> AsyncGenerator[dict[str, object], None]:
cm: Callable[
[Starlette, svcs.Registry], contextlib.AbstractAsyncContextManager
]
if inspect.isasyncgenfunction(self._lifespan):
cm = contextlib.asynccontextmanager(self._lifespan)
else:
cm = self._lifespan # type: ignore[assignment] # ty: ignore[invalid-assignment]
# When lifespans are composed (e.g. FastAPI's router lifespans),
# the first svcs lifespan wins and detaches the registry on exit.
owns_app_state = not hasattr(app.state, _KEY_REGISTRY)
if owns_app_state:
setattr(app.state, _KEY_REGISTRY, self.registry)
try:
async with self.registry, cm(app, self.registry) as state:
self._state = state or {}
self._state[_KEY_REGISTRY] = self.registry
yield self._state
finally:
if owns_app_state:
delattr(app.state, _KEY_REGISTRY)
[docs]
def get_registry(app: Starlette | TestClient) -> svcs.Registry:
"""
Get the registry that :class:`lifespan` has attached to *app*.
The registry is attached when the application starts, so this only works
on a running application.
Args:
app:
A Starlette application with a *svcs*-aware lifespan, or a
:class:`starlette.testclient.TestClient` wrapping one.
Raises:
LookupError: If no registry is attached to *app*.
.. versionadded:: 26.2.0
"""
try:
if not isinstance(app, Starlette):
app = cast("Starlette", app.app)
return getattr(app.state, _KEY_REGISTRY) # type: ignore[no-any-return]
except AttributeError:
msg = "No svcs registry on app."
raise LookupError(msg) from None
[docs]
@attrs.define
class SVCSMiddleware:
"""
Attach a :class:`svcs.Container` to the request state, based on a registry
that has been put on the request state by :class:`lifespan`. Closes the
container at the end of a request or websocket connection.
"""
app: ASGIApp
async def __call__(
self, scope: Scope, receive: Receive, send: Send
) -> None:
if scope["type"] not in ("http", "websocket"):
return await self.app(scope, receive, send)
async with svcs.Container(scope["state"][_KEY_REGISTRY]) as con:
scope["state"][_KEY_CONTAINER] = con
return await self.app(scope, receive, send)
[docs]
def get_pings(request: Request) -> list[svcs.ServicePing]:
"""
Same as :meth:`svcs.Container.get_pings`, but uses the container from
*request*.
See Also:
:ref:`aiohttp-health`
"""
return svcs_from(request).get_pings()
[docs]
async def aget_abstract(request: Request, *svc_types: _ServiceType) -> Any:
"""
Same as :meth:`svcs.Container.aget_abstract()`, but uses container from
*request*.
.. deprecated:: 26.1.0
"""
return await svcs_from(request).aget_abstract(*svc_types)
@overload
async def aget(request: Request, svc_type: TypeForm[T1], /) -> T1: ...
@overload
async def aget(
request: Request,
svc_type1: TypeForm[T1],
svc_type2: TypeForm[T2],
/,
) -> tuple[T1, T2]: ...
@overload
async def aget(
request: Request,
svc_type1: TypeForm[T1],
svc_type2: TypeForm[T2],
svc_type3: TypeForm[T3],
/,
) -> tuple[T1, T2, T3]: ...
@overload
async def aget(
request: Request,
svc_type1: TypeForm[T1],
svc_type2: TypeForm[T2],
svc_type3: TypeForm[T3],
svc_type4: TypeForm[T4],
/,
) -> tuple[T1, T2, T3, T4]: ...
@overload
async def aget(
request: Request,
svc_type1: TypeForm[T1],
svc_type2: TypeForm[T2],
svc_type3: TypeForm[T3],
svc_type4: TypeForm[T4],
svc_type5: TypeForm[T5],
/,
) -> tuple[T1, T2, T3, T4, T5]: ...
@overload
async def aget(
request: Request,
svc_type1: TypeForm[T1],
svc_type2: TypeForm[T2],
svc_type3: TypeForm[T3],
svc_type4: TypeForm[T4],
svc_type5: TypeForm[T5],
svc_type6: TypeForm[T6],
/,
) -> tuple[T1, T2, T3, T4, T5, T6]: ...
@overload
async def aget(
request: Request,
svc_type1: TypeForm[T1],
svc_type2: TypeForm[T2],
svc_type3: TypeForm[T3],
svc_type4: TypeForm[T4],
svc_type5: TypeForm[T5],
svc_type6: TypeForm[T6],
svc_type7: TypeForm[T7],
/,
) -> tuple[T1, T2, T3, T4, T5, T6, T7]: ...
@overload
async def aget(
request: Request,
svc_type1: TypeForm[T1],
svc_type2: TypeForm[T2],
svc_type3: TypeForm[T3],
svc_type4: TypeForm[T4],
svc_type5: TypeForm[T5],
svc_type6: TypeForm[T6],
svc_type7: TypeForm[T7],
svc_type8: TypeForm[T8],
/,
) -> tuple[T1, T2, T3, T4, T5, T6, T7, T8]: ...
@overload
async def aget(
request: Request,
svc_type1: TypeForm[T1],
svc_type2: TypeForm[T2],
svc_type3: TypeForm[T3],
svc_type4: TypeForm[T4],
svc_type5: TypeForm[T5],
svc_type6: TypeForm[T6],
svc_type7: TypeForm[T7],
svc_type8: TypeForm[T8],
svc_type9: TypeForm[T9],
/,
) -> tuple[T1, T2, T3, T4, T5, T6, T7, T8, T9]: ...
@overload
async def aget(
request: Request,
svc_type1: TypeForm[T1],
svc_type2: TypeForm[T2],
svc_type3: TypeForm[T3],
svc_type4: TypeForm[T4],
svc_type5: TypeForm[T5],
svc_type6: TypeForm[T6],
svc_type7: TypeForm[T7],
svc_type8: TypeForm[T8],
svc_type9: TypeForm[T9],
svc_type10: TypeForm[T10],
/,
) -> tuple[T1, T2, T3, T4, T5, T6, T7, T8, T9, T10]: ...
[docs]
async def aget(request: Request, *svc_types: _ServiceType) -> object:
"""
Same as :meth:`svcs.Container.aget`, but uses the container from *request*.
"""
return await svcs_from(request).aget(*svc_types)