|
1 | 1 | """Starlette/FastAPI context middleware.""" |
2 | 2 |
|
3 | | -from typing import Any |
| 3 | +from typing import Any, List, Optional |
4 | 4 |
|
5 | | -from sap_cloud_sdk.core.runtime_context._context import async_sdk_context |
| 5 | +from sap_cloud_sdk.core.runtime_context._context import RequestContext, async_sdk_context |
| 6 | +from sap_cloud_sdk.core.runtime_context._envelope import RequestEnvelope |
6 | 7 | from sap_cloud_sdk.core.runtime_context._protocol import ContextProvider |
7 | 8 |
|
8 | 9 | try: |
|
16 | 17 | ) from exc |
17 | 18 |
|
18 | 19 |
|
| 20 | +def _merge(contexts: List[RequestContext]) -> RequestContext: |
| 21 | + """Merge multiple RequestContexts — first non-None value wins per field.""" |
| 22 | + merged = RequestContext() |
| 23 | + for ctx in contexts: |
| 24 | + if merged.tenant_id is None: |
| 25 | + merged.tenant_id = ctx.tenant_id |
| 26 | + if merged.user_id is None: |
| 27 | + merged.user_id = ctx.user_id |
| 28 | + if merged.trigger_type is None: |
| 29 | + merged.trigger_type = ctx.trigger_type |
| 30 | + merged.extras.update(ctx.extras) |
| 31 | + return merged |
| 32 | + |
| 33 | + |
19 | 34 | class StarletteContextMiddleware(BaseHTTPMiddleware): |
20 | 35 | """Starlette/FastAPI middleware that populates the SDK runtime context. |
21 | 36 |
|
22 | | - Runs *provider*.extract() on every inbound request and makes the result |
23 | | - available via :func:`~sap_cloud_sdk.core.runtime_context.get_context` |
24 | | - for the duration of that request. |
| 37 | + Builds a :class:`~sap_cloud_sdk.core.runtime_context.RequestEnvelope` from |
| 38 | + each inbound request, runs all *providers* against it, and merges the results |
| 39 | + into a single :class:`~sap_cloud_sdk.core.runtime_context.RequestContext` |
| 40 | + available via :func:`~sap_cloud_sdk.core.runtime_context.get_context` for |
| 41 | + the duration of that request. |
| 42 | +
|
| 43 | + First non-None value wins per field when merging. Extras are union-merged |
| 44 | + (later providers can add keys, but not overwrite earlier ones). |
25 | 45 |
|
26 | 46 | Usage:: |
27 | 47 |
|
28 | | - from starlette.applications import Starlette |
29 | 48 | from sap_cloud_sdk import bootstrap |
30 | 49 |
|
31 | | - app = Starlette(...) |
32 | | - bootstrap(app) |
33 | | -
|
34 | | - Or manually:: |
35 | | -
|
36 | | - from sap_cloud_sdk.core.runtime_context.starlette import StarletteContextMiddleware |
37 | | - from sap_cloud_sdk.core.runtime_context import IASContextProvider |
| 50 | + bootstrap(app) # IASContextProvider by default |
38 | 51 |
|
39 | | - app.add_middleware(StarletteContextMiddleware, provider=IASContextProvider()) |
| 52 | + # or with multiple providers: |
| 53 | + bootstrap(app, providers=[IASContextProvider(), MyCustomProvider()]) |
40 | 54 | """ |
41 | 55 |
|
42 | | - def __init__(self, app: Any, provider: ContextProvider) -> None: |
| 56 | + def __init__(self, app: Any, providers: List[ContextProvider]) -> None: |
43 | 57 | super().__init__(app) |
44 | | - self._provider = provider |
| 58 | + self._providers = providers |
45 | 59 |
|
46 | 60 | async def dispatch(self, request: Request, call_next: Any) -> Response: |
47 | | - ctx = self._provider.extract(request) |
| 61 | + envelope = RequestEnvelope(headers=dict(request.headers)) |
| 62 | + ctx = _merge([p.extract(envelope) for p in self._providers]) |
48 | 63 | async with async_sdk_context(ctx): |
49 | 64 | return await call_next(request) |
0 commit comments