diff --git a/pyproject.toml b/pyproject.toml index 94264213..76f92cb0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "sap-cloud-sdk" -version = "0.37.0" +version = "0.38.0" description = "SAP Cloud SDK for Python" readme = "README.md" license = "Apache-2.0" diff --git a/src/sap_cloud_sdk/core/auditlog_ng/client.py b/src/sap_cloud_sdk/core/auditlog_ng/client.py index 258e7960..12748255 100644 --- a/src/sap_cloud_sdk/core/auditlog_ng/client.py +++ b/src/sap_cloud_sdk/core/auditlog_ng/client.py @@ -7,7 +7,8 @@ import json import os import uuid -from typing import Optional +from datetime import datetime, timezone +from typing import Any, Optional, cast import protovalidate from protovalidate import ValidationError as ProtoValidationError @@ -37,6 +38,28 @@ from sap_cloud_sdk.core.auditlog_ng.exceptions import ValidationError from sap_cloud_sdk.core.telemetry import Module, Operation, record_metrics from sap_cloud_sdk.core.telemetry.config import ENV_OTLP_PROTOCOL +from sap_cloud_sdk.ias._context import get_auth_context + + +def _fill_common_from_auth_context(event: Message) -> None: + """Back-fill common fields from the request auth context. + + Accesses event.common directly — protobuf creates the sub-message on first + mutation, so callers that never touched common still get it populated. + Sets tenant_id and user_initiator_id from IAS claims (if present and not + already set), and sets timestamp to now if the caller left it at zero. + """ + if not hasattr(event, "common"): + return + common = cast(Any, event.common) + claims = get_auth_context() + if claims is not None: + if claims.app_tid and not common.tenant_id: + common.tenant_id = claims.app_tid + if claims.user_uuid and not common.user_initiator_id: + common.user_initiator_id = claims.user_uuid + if common.timestamp.seconds == 0: + common.timestamp.FromDatetime(datetime.now(timezone.utc)) def _create_log_exporter( @@ -178,6 +201,8 @@ def send( if format not in {"protobuf-binary", "json"}: raise ValueError("format must be 'protobuf-binary' or 'json'") + _fill_common_from_auth_context(event) + try: protovalidate.validate(event) except ProtoValidationError as e: diff --git a/src/sap_cloud_sdk/core/auditlog_ng/user-guide.md b/src/sap_cloud_sdk/core/auditlog_ng/user-guide.md index f6d58427..5532efb9 100644 --- a/src/sap_cloud_sdk/core/auditlog_ng/user-guide.md +++ b/src/sap_cloud_sdk/core/auditlog_ng/user-guide.md @@ -168,6 +168,8 @@ event.object_type = "resource" event.object_id = "resource-001" ``` +> **Tip:** When using `StarletteIASTelemetryMiddleware` (see [Automatic tenant and user injection](#automatic-tenant-and-user-injection)), `common.tenant_id` and `common.user_initiator_id` are filled automatically from the incoming IAS JWT. You only need to set them explicitly if you want to override the values from the token. + ### Step 4: Send the Event **Binary protobuf:** @@ -284,6 +286,55 @@ Events are validated against protobuf constraints using `protovalidate` before s --- +## Automatic tenant and user injection + +When `StarletteIASTelemetryMiddleware` is registered on your app, it parses the +incoming `Authorization: Bearer ` header on every request and stores the +IAS claims in the current async context. + +`AuditClient.send()` reads that context automatically before validation and +back-fills two fields on the event's `common` block — only if they are not +already set by the caller: + +| Field populated | IAS claim used | +|---|---| +| `common.tenant_id` | `app_tid` | +| `common.user_initiator_id` | `user_uuid` | + +### Setup + +Register the middleware once when your app starts: + +```python +from sap_cloud_sdk.core.telemetry import auto_instrument +from sap_cloud_sdk.core.telemetry.middleware import StarletteIASTelemetryMiddleware + +app = FastAPI(...) +auto_instrument(middlewares=[StarletteIASTelemetryMiddleware(app=app)]) +``` + +### Usage + +With the middleware in place, you can omit `tenant_id` and `user_initiator_id` +from every event — they are injected automatically: + +```python +event = pb.DataAccess() +event.common.timestamp.FromDatetime(datetime.now(timezone.utc)) +# tenant_id and user_initiator_id are filled from the IAS JWT automatically +event.channel_type = "API" +event.channel_id = "agent-v1" +event.object_type = "resource" +event.object_id = "resource-001" + +event_id = client.send(event) +``` + +If neither the middleware nor an explicit value provides `tenant_id`, the event +will fail `protovalidate` validation and raise a `ValidationError`. + +--- + ## Running the Unit Tests ```bash diff --git a/src/sap_cloud_sdk/core/telemetry/middleware/starlette_a2a.py b/src/sap_cloud_sdk/core/telemetry/middleware/starlette_a2a.py index 1b34c733..3bbde413 100644 --- a/src/sap_cloud_sdk/core/telemetry/middleware/starlette_a2a.py +++ b/src/sap_cloud_sdk/core/telemetry/middleware/starlette_a2a.py @@ -10,7 +10,7 @@ ATTR_USER_ID, ) from sap_cloud_sdk.core.telemetry.middleware.base import TelemetryMiddleware -from sap_cloud_sdk.ias import parse_token +from sap_cloud_sdk.ias import parse_token, set_auth_context try: from starlette.middleware.base import BaseHTTPMiddleware @@ -31,7 +31,9 @@ def __init__(self, app: Any, attrs_var: ContextVar[Dict[str, Any]]) -> None: self._attrs_var = attrs_var async def dispatch(self, request: Request, call_next: Any) -> Response: - token = self._attrs_var.set(_extract_ias_attrs(request)) + claims, attrs = _parse_request(request) + set_auth_context(claims) + token = self._attrs_var.set(attrs) try: return await call_next(request) finally: @@ -78,23 +80,22 @@ def get_attributes(self) -> Dict[str, Any]: return self._attrs_var.get() -def _extract_ias_attrs(request: Request) -> Dict[str, Any]: - """Parse the Authorization header and return telemetry attributes.""" +def _parse_request(request: Request): + """Parse the Authorization header and return (IASClaims, telemetry_attrs).""" auth = request.headers.get("authorization", "") - if not auth: - return {} - try: - claims = parse_token(auth) - except Exception as e: - logger.debug("IAS token parsing failed, skipping telemetry attrs: %s", e) - return {} - + claims = None attrs: Dict[str, Any] = {} - if claims.sap_gtid: - attrs[ATTR_SAP_TENANT_ID] = claims.sap_gtid - if claims.user_uuid: - attrs[ATTR_USER_ID] = claims.user_uuid + if auth: + try: + claims = parse_token(auth) + except Exception as e: + logger.debug("IAS token parsing failed, skipping telemetry attrs: %s", e) + if claims is not None: + if claims.sap_gtid: + attrs[ATTR_SAP_TENANT_ID] = claims.sap_gtid + if claims.user_uuid: + attrs[ATTR_USER_ID] = claims.user_uuid origin = request.headers.get("x-sap-origin") if origin: attrs[ATTR_SAP_TRIGGER_TYPE] = origin - return attrs + return claims, attrs diff --git a/src/sap_cloud_sdk/ias/__init__.py b/src/sap_cloud_sdk/ias/__init__.py index 88df1774..c2459702 100644 --- a/src/sap_cloud_sdk/ias/__init__.py +++ b/src/sap_cloud_sdk/ias/__init__.py @@ -12,11 +12,14 @@ print(claims.email) # user email (when email scope requested) """ +from sap_cloud_sdk.ias._context import get_auth_context, set_auth_context from sap_cloud_sdk.ias._token import IASClaims, parse_token from sap_cloud_sdk.ias.exceptions import IASTokenError __all__ = [ "IASClaims", - "parse_token", "IASTokenError", + "get_auth_context", + "parse_token", + "set_auth_context", ] diff --git a/src/sap_cloud_sdk/ias/_context.py b/src/sap_cloud_sdk/ias/_context.py new file mode 100644 index 00000000..ff578401 --- /dev/null +++ b/src/sap_cloud_sdk/ias/_context.py @@ -0,0 +1,20 @@ +"""Per-request IAS auth context via ContextVar.""" + +from contextvars import ContextVar +from typing import Optional + +from sap_cloud_sdk.ias._token import IASClaims + +_auth_context_var: ContextVar[Optional[IASClaims]] = ContextVar( + "ias_auth_context", default=None +) + + +def set_auth_context(claims: Optional[IASClaims]) -> None: + """Store IAS claims for the current async context.""" + _auth_context_var.set(claims) + + +def get_auth_context() -> Optional[IASClaims]: + """Return the IAS claims set for the current async context, or None.""" + return _auth_context_var.get() diff --git a/tests/core/unit/auditlog_ng/unit/test_client.py b/tests/core/unit/auditlog_ng/unit/test_client.py index 238d8e5f..8d22a882 100644 --- a/tests/core/unit/auditlog_ng/unit/test_client.py +++ b/tests/core/unit/auditlog_ng/unit/test_client.py @@ -2,15 +2,18 @@ from __future__ import annotations +from datetime import datetime, timezone from typing import TypedDict, Unpack from unittest.mock import MagicMock, Mock, patch import pytest -from sap_cloud_sdk.core.auditlog_ng.client import AuditClient +from sap_cloud_sdk.core.auditlog_ng.client import AuditClient, _fill_common_from_auth_context from sap_cloud_sdk.core.auditlog_ng.config import AuditLogNGConfig, SCHEMA_URL from sap_cloud_sdk.core.auditlog_ng.exceptions import ValidationError +from sap_cloud_sdk.core.auditlog_ng.gen.sap.auditlog.auditevent.v2 import auditevent_pb2 as pb from sap_cloud_sdk.core.telemetry import Module, Operation +from sap_cloud_sdk.ias._context import set_auth_context class ConfigKwargs(TypedDict, total=False): @@ -335,3 +338,57 @@ def test_unsupported_protocol_raises( monkeypatch.setenv("OTEL_EXPORTER_OTLP_PROTOCOL", "http/json") with pytest.raises(ValueError, match="Unsupported OTEL_EXPORTER_OTLP_PROTOCOL"): AuditClient(_make_config()) + + +class TestFillCommonFromAuthContext: + def setup_method(self): + set_auth_context(None) + + def teardown_method(self): + set_auth_context(None) + + def _make_claims(self, app_tid="tid-123", user_uuid="uid-456"): + claims = MagicMock() + claims.app_tid = app_tid + claims.user_uuid = user_uuid + return claims + + def test_no_context_sets_timestamp_only(self): + event = pb.DataAccess() + _fill_common_from_auth_context(event) + assert event.common.timestamp.seconds > 0 + assert event.common.tenant_id == "" + assert event.common.user_initiator_id == "" + + def test_with_context_fills_tenant_and_user(self): + set_auth_context(self._make_claims()) + event = pb.DataAccess() + _fill_common_from_auth_context(event) + assert event.common.tenant_id == "tid-123" + assert event.common.user_initiator_id == "uid-456" + assert event.common.timestamp.seconds > 0 + + def test_does_not_overwrite_explicit_tenant(self): + set_auth_context(self._make_claims(app_tid="ctx-tid")) + event = pb.DataAccess() + event.common.tenant_id = "explicit-tid" + _fill_common_from_auth_context(event) + assert event.common.tenant_id == "explicit-tid" + + def test_does_not_overwrite_explicit_user(self): + set_auth_context(self._make_claims(user_uuid="ctx-uid")) + event = pb.DataAccess() + event.common.user_initiator_id = "explicit-uid" + _fill_common_from_auth_context(event) + assert event.common.user_initiator_id == "explicit-uid" + + def test_does_not_overwrite_explicit_timestamp(self): + fixed = datetime(2024, 1, 1, tzinfo=timezone.utc) + event = pb.DataAccess() + event.common.timestamp.FromDatetime(fixed) + _fill_common_from_auth_context(event) + assert event.common.timestamp.seconds == int(fixed.timestamp()) + + def test_event_without_common_field_is_skipped(self): + event = MagicMock(spec=[]) # no 'common' attribute + _fill_common_from_auth_context(event) # must not raise diff --git a/tests/core/unit/telemetry/middleware/test_starlette_a2a.py b/tests/core/unit/telemetry/middleware/test_starlette_a2a.py index e0cfdda3..871798d9 100644 --- a/tests/core/unit/telemetry/middleware/test_starlette_a2a.py +++ b/tests/core/unit/telemetry/middleware/test_starlette_a2a.py @@ -6,7 +6,7 @@ from sap_cloud_sdk.core.telemetry.constants import ATTR_SAP_TRIGGER_TYPE, ATTR_SAP_TENANT_ID, ATTR_USER_ID from sap_cloud_sdk.core.telemetry.middleware.starlette_a2a import ( StarletteIASTelemetryMiddleware, - _extract_ias_attrs, + _parse_request, ) _PATCH_PARSE = "sap_cloud_sdk.core.telemetry.middleware.starlette_a2a.parse_token" @@ -55,19 +55,19 @@ def test_two_instances_do_not_interfere(self): mw2._attrs_var.reset(t2) -class TestExtractIasAttrs: +class TestParseRequest: def test_extracts_tenant_and_user(self): claims = _make_claims(sap_gtid="t1", user_uuid="u1") request = _make_request({"authorization": "Bearer tok"}) with patch(_PATCH_PARSE, return_value=claims): - result = _extract_ias_attrs(request) + _, result = _parse_request(request) assert result == {ATTR_SAP_TENANT_ID: "t1", ATTR_USER_ID: "u1"} def test_omits_missing_tenant(self): claims = _make_claims(sap_gtid=None, user_uuid="u1") request = _make_request({"authorization": "Bearer tok"}) with patch(_PATCH_PARSE, return_value=claims): - result = _extract_ias_attrs(request) + _, result = _parse_request(request) assert result == {ATTR_USER_ID: "u1"} assert ATTR_SAP_TENANT_ID not in result @@ -75,44 +75,53 @@ def test_omits_missing_user(self): claims = _make_claims(sap_gtid="t1", user_uuid=None) request = _make_request({"authorization": "Bearer tok"}) with patch(_PATCH_PARSE, return_value=claims): - result = _extract_ias_attrs(request) + _, result = _parse_request(request) assert result == {ATTR_SAP_TENANT_ID: "t1"} assert ATTR_USER_ID not in result def test_returns_empty_when_no_auth_header(self): request = _make_request({}) with patch(_PATCH_PARSE) as mock_parse: - result = _extract_ias_attrs(request) + parsed_claims, result = _parse_request(request) mock_parse.assert_not_called() assert result == {} + assert parsed_claims is None def test_returns_empty_on_parse_error(self): request = _make_request({"authorization": "Bearer bad"}) with patch(_PATCH_PARSE, side_effect=ValueError("bad token")): - result = _extract_ias_attrs(request) + parsed_claims, result = _parse_request(request) assert result == {} + assert parsed_claims is None def test_returns_empty_when_both_claims_absent(self): claims = _make_claims(sap_gtid=None, user_uuid=None) request = _make_request({"authorization": "Bearer tok"}) with patch(_PATCH_PARSE, return_value=claims): - result = _extract_ias_attrs(request) + _, result = _parse_request(request) assert result == {} def test_includes_origin_header_when_present(self): claims = _make_claims(sap_gtid="t1", user_uuid="u1") request = _make_request({"authorization": "Bearer tok", "x-sap-origin": "ui5"}) with patch(_PATCH_PARSE, return_value=claims): - result = _extract_ias_attrs(request) + _, result = _parse_request(request) assert result[ATTR_SAP_TRIGGER_TYPE] == "ui5" def test_omits_origin_attr_when_header_absent(self): claims = _make_claims(sap_gtid="t1", user_uuid="u1") request = _make_request({"authorization": "Bearer tok"}) with patch(_PATCH_PARSE, return_value=claims): - result = _extract_ias_attrs(request) + _, result = _parse_request(request) assert ATTR_SAP_TRIGGER_TYPE not in result + def test_returns_claims_object(self): + claims = _make_claims(sap_gtid="t1", user_uuid="u1") + request = _make_request({"authorization": "Bearer tok"}) + with patch(_PATCH_PARSE, return_value=claims): + parsed_claims, _ = _parse_request(request) + assert parsed_claims is claims + class TestInnerMiddlewareDispatch: def _get_inner_class_and_kwargs(self, mw: StarletteIASTelemetryMiddleware):