diff --git a/ddtrace/__init__.py b/ddtrace/__init__.py index 941c8ac31d9..cf312bb8616 100644 --- a/ddtrace/__init__.py +++ b/ddtrace/__init__.py @@ -12,7 +12,8 @@ configure_ddtrace_logger() # noqa: E402 # Enable telemetry writer and excepthook as early as possible to ensure we capture any exceptions from initialization -import ddtrace.internal.telemetry # noqa: E402, F401 +from ddtrace.internal.serverless import in_aws_lambda_microvm # noqa: E402 +import ddtrace.internal.telemetry # noqa: F401,E402 from ._monkey import patch # noqa: E402 from ._monkey import patch_all # noqa: E402 @@ -25,6 +26,13 @@ from .version import __version__ # noqa: E402 +# Register import-time hooks that depend on ddtrace.config being exported. +if in_aws_lambda_microvm(): + from ddtrace.internal.runtime import listen_for_identity_refresh_hooks # noqa: E402,I001 + from .internal import core as _core # noqa: E402 + + listen_for_identity_refresh_hooks(_core.on) + # TODO: Deprecate accessing tracer from ddtrace.__init__ module in v4.0 if env.get("_DD_GLOBAL_TRACER_INIT", "true").lower() in ("1", "true"): from ddtrace.trace import tracer # noqa: F401 diff --git a/ddtrace/internal/_runtime_id.py b/ddtrace/internal/_runtime_id.py index fe4472ee2b6..051e6faf53e 100644 --- a/ddtrace/internal/_runtime_id.py +++ b/ddtrace/internal/_runtime_id.py @@ -2,6 +2,10 @@ import typing as t import uuid +from ddtrace.internal.constants import WEB_REQUEST_STARTING_EVENT +from ddtrace.internal.serverless import MICROVM_RUN_HOOK_METHOD +from ddtrace.internal.serverless import MICROVM_RUN_HOOK_PATH +from ddtrace.internal.serverless import in_aws_lambda_microvm from ddtrace.internal.settings import env from . import forksafe @@ -16,6 +20,8 @@ "get_runtime_id", "get_parent_runtime_id", "get_runtime_propagation_envs", + "listen_for_identity_refresh_hooks", + "maybe_refresh_identity", "refresh_identity", ] @@ -37,6 +43,9 @@ def _generate_runtime_id() -> str: # Module-level set[...] in Python 3.10 affects import timing. See packages.py for details. _ON_RUNTIME_ID_CHANGE: t.Set[t.Callable[[str], None]] = set() # noqa: UP006 _ON_RUNTIME_IDENTITY_REFRESH: t.Set[t.Callable[[str], None]] = set() # noqa: UP006 +# MicroVM refreshes share this lock with consumers that must not observe a partially refreshed +# identity. Non-MicroVM callers do not acquire it. +_RUNTIME_IDENTITY_REFRESH_LOCK = forksafe.RLock() def on_runtime_id_change(cb: t.Callable[[str], None]) -> None: @@ -58,6 +67,11 @@ def on_runtime_identity_refresh(cb: t.Callable[[str], None]) -> None: _ON_RUNTIME_IDENTITY_REFRESH.add(cb) +def get_runtime_identity_refresh_lock() -> t.ContextManager[None]: + """Return the lock that serializes a MicroVM identity refresh with its consumers.""" + return t.cast(t.ContextManager[None], _RUNTIME_IDENTITY_REFRESH_LOCK) + + def _notify_runtime_id_callbacks(callbacks: t.Set[t.Callable[[str], None]]) -> None: # noqa: UP006 for cb in list(callbacks): try: @@ -66,18 +80,44 @@ def _notify_runtime_id_callbacks(callbacks: t.Set[t.Callable[[str], None]]) -> N log.exception("Exception ignored in runtime ID callback %r", cb) -def _notify_runtime_identity_refresh_callbacks(*, raise_on_error: bool = False) -> None: # noqa: UP006 +def _notify_runtime_identity_refresh_callbacks( + *, + callbacks: t.Optional[t.List[t.Callable[[str], None]]] = None, # noqa: UP006 + raise_on_error: bool = False, +) -> None: # noqa: UP006 # Direct refresh callers keep subscriber failures isolated so every component gets # a chance to rebuild. The MicroVM coordinator opts into propagation so a failed # rebuild leaves its completion guard unset and the same identity can be retried. - for cb in list(_ON_RUNTIME_IDENTITY_REFRESH): - if raise_on_error: - cb(_RUNTIME_ID) - continue + # A caller-supplied list is the MicroVM transition's pending work queue. The + # permanent registry must remain unchanged so later refreshes still notify all + # registered callbacks. + # Only the MicroVM coordinator supplies a pending queue. It runs every pending callback + # before propagating, so a callback that keeps failing cannot starve the ones behind it + # on each retry. + defer_errors = raise_on_error and callbacks is not None + if callbacks is None: + callbacks = list(_ON_RUNTIME_IDENTITY_REFRESH) + + first_error: t.Optional[Exception] = None + for cb in list(callbacks): try: cb(_RUNTIME_ID) - except Exception: - log.exception("Exception ignored in runtime ID callback %r", cb) + except Exception as e: + if not raise_on_error: + log.exception("Exception ignored in runtime ID callback %r", cb) + continue + if not defer_errors: + raise + if first_error is None: + first_error = e + log.exception("Runtime ID callback %r failed and stays pending", cb) + continue + + # Removing only after return records completion; a raised callback stays pending. + callbacks.remove(cb) + + if first_error is not None: + raise first_error def _refresh_runtime_id() -> None: @@ -99,7 +139,11 @@ def _set_runtime_id() -> None: _refresh_runtime_id() -def refresh_identity(raise_on_error: bool = False) -> None: +def refresh_identity( + raise_on_error: bool = False, + *, + _callback_snapshot: t.Optional[t.List[t.Callable[[str], None]]] = None, # noqa: UP006 +) -> None: # noqa: UP006 """Regenerate the runtime ID without recording fork lineage. Unlike a fork, this does not update _PARENT_RUNTIME_ID / _ANCESTOR_RUNTIME_ID: @@ -113,8 +157,88 @@ def refresh_identity(raise_on_error: bool = False) -> None: # Notify consumers that only need the new ID first. The explicit refresh # callbacks below are for components that must rebuild restore-sensitive # state, which is different from the fork handling in _set_runtime_id(). + if in_aws_lambda_microvm(): + with _RUNTIME_IDENTITY_REFRESH_LOCK: + _refresh_identity(raise_on_error, _callback_snapshot) + else: + _refresh_identity(raise_on_error, _callback_snapshot) + + +def _refresh_identity( + raise_on_error: bool, + callback_snapshot: t.Optional[t.List[t.Callable[[str], None]]], # noqa: UP006 +) -> None: # noqa: UP006 _refresh_runtime_id() - _notify_runtime_identity_refresh_callbacks(raise_on_error=raise_on_error) + if callback_snapshot is not None: + # Replace the caller-owned queue after rotation. This preserves the normal + # refresh ordering without duplicating callbacks if the list was reused. + callback_snapshot[:] = _ON_RUNTIME_IDENTITY_REFRESH + + _notify_runtime_identity_refresh_callbacks(callbacks=callback_snapshot, raise_on_error=raise_on_error) + + +# Multiple request layers can observe the same /run hook. Refresh identity once per +# process so a single logical MicroVM instance gets one runtime-id rotation. +_IDENTITY_REFRESH_HOOK_REFRESHED = forksafe.Event() +_IDENTITY_REFRESH_HOOK_REFRESH_LOCK = forksafe.Lock() +# Keep a failed transition retryable without rotating the identity again. +_IDENTITY_REFRESH_HOOK_RUNTIME_ID: t.Optional[str] = None +# This is per-transition state, unlike _ON_RUNTIME_IDENTITY_REFRESH. Callbacks are +# removed here only after succeeding for _IDENTITY_REFRESH_HOOK_RUNTIME_ID. +_IDENTITY_REFRESH_HOOK_PENDING_CALLBACKS: t.Optional[t.List[t.Callable[[str], None]]] = None # noqa: UP006 + + +def listen_for_identity_refresh_hooks( + on_event: t.Callable[[str, t.Callable[[t.Optional[str], t.Optional[str]], None]], None], +) -> None: + """Refresh MicroVM identity from request events emitted before root span creation.""" + if not in_aws_lambda_microvm(): + return + + on_event(WEB_REQUEST_STARTING_EVENT, maybe_refresh_identity) + + +def maybe_refresh_identity(method: t.Optional[str], path: t.Optional[str]) -> None: + """Call refresh_identity() if this request is the AWS Lambda MicroVM /run hook.""" + if not in_aws_lambda_microvm(): + return + if not method or not path: + return + if method != MICROVM_RUN_HOOK_METHOD or path != MICROVM_RUN_HOOK_PATH: + return + + global _IDENTITY_REFRESH_HOOK_PENDING_CALLBACKS, _IDENTITY_REFRESH_HOOK_RUNTIME_ID + with _IDENTITY_REFRESH_HOOK_REFRESH_LOCK: + if _IDENTITY_REFRESH_HOOK_REFRESHED.is_set(): + return + + # Rotate once per transition; after a callback failure, retry only callbacks that + # did not complete against the current ID. The completion guard is set only after + # every callback succeeds. + if _IDENTITY_REFRESH_HOOK_RUNTIME_ID != _RUNTIME_ID: + _IDENTITY_REFRESH_HOOK_PENDING_CALLBACKS = [] + # Keep refresh_identity() as the single refresh entry point. Its private + # snapshot sink lets this hook retain progress when a subscriber fails. + try: + refresh_identity(raise_on_error=True, _callback_snapshot=_IDENTITY_REFRESH_HOOK_PENDING_CALLBACKS) + except Exception: + # The ID has already rotated before invoking explicit subscribers. Retry those + # subscribers against that ID. + _IDENTITY_REFRESH_HOOK_RUNTIME_ID = _RUNTIME_ID + raise + _IDENTITY_REFRESH_HOOK_RUNTIME_ID = _RUNTIME_ID + else: + # refresh_identity() releases this lock before we retry only the + # callbacks that previously failed. Keep the retry serialized with + # consumers of the refreshed identity as well. + with _RUNTIME_IDENTITY_REFRESH_LOCK: + _notify_runtime_identity_refresh_callbacks( + callbacks=_IDENTITY_REFRESH_HOOK_PENDING_CALLBACKS, + raise_on_error=True, + ) + + _IDENTITY_REFRESH_HOOK_REFRESHED.set() + _IDENTITY_REFRESH_HOOK_PENDING_CALLBACKS = None def get_runtime_id() -> str: diff --git a/ddtrace/internal/constants.py b/ddtrace/internal/constants.py index 2cbd3d91161..94bb07f5a59 100644 --- a/ddtrace/internal/constants.py +++ b/ddtrace/internal/constants.py @@ -8,6 +8,7 @@ from ddtrace.constants import USER_REJECT +WEB_REQUEST_STARTING_EVENT = "web.request.starting" _WEB_REQUEST_STARTING_DISPATCHED = "_ddtrace_web_request_starting_dispatched" diff --git a/ddtrace/internal/runtime/__init__.py b/ddtrace/internal/runtime/__init__.py index 7657b2c46a5..489bb5dd51b 100644 --- a/ddtrace/internal/runtime/__init__.py +++ b/ddtrace/internal/runtime/__init__.py @@ -2,7 +2,10 @@ from ddtrace.internal._runtime_id import get_parent_runtime_id from ddtrace.internal._runtime_id import get_process_role from ddtrace.internal._runtime_id import get_runtime_id +from ddtrace.internal._runtime_id import get_runtime_identity_refresh_lock from ddtrace.internal._runtime_id import get_runtime_propagation_envs +from ddtrace.internal._runtime_id import listen_for_identity_refresh_hooks +from ddtrace.internal._runtime_id import maybe_refresh_identity from ddtrace.internal._runtime_id import on_runtime_id_change from ddtrace.internal._runtime_id import on_runtime_identity_refresh from ddtrace.internal._runtime_id import refresh_identity @@ -12,9 +15,12 @@ "get_ancestor_runtime_id", "get_process_role", "get_runtime_id", + "get_runtime_identity_refresh_lock", "get_parent_runtime_id", "get_runtime_propagation_envs", "on_runtime_id_change", "on_runtime_identity_refresh", + "listen_for_identity_refresh_hooks", + "maybe_refresh_identity", "refresh_identity", ] diff --git a/ddtrace/internal/serverless/__init__.py b/ddtrace/internal/serverless/__init__.py index d9017c9dbbe..41ff40671cd 100644 --- a/ddtrace/internal/serverless/__init__.py +++ b/ddtrace/internal/serverless/__init__.py @@ -3,6 +3,11 @@ from ddtrace.internal.settings import env +# Fixed platform request details for the AWS Lambda MicroVM /run lifecycle hook. +MICROVM_RUN_HOOK_METHOD = "POST" +MICROVM_RUN_HOOK_PATH = "/aws/lambda-microvms/runtime/v1/run" + + def in_aws_lambda() -> bool: """Returns whether the environment is an AWS Lambda. This is accomplished by checking if the AWS_LAMBDA_FUNCTION_NAME environment diff --git a/releasenotes/notes/aws-lambda-microvm-identity-refresh-3a672cd6bcbad16d.yaml b/releasenotes/notes/aws-lambda-microvm-identity-refresh-3a672cd6bcbad16d.yaml new file mode 100644 index 00000000000..928b291dd51 --- /dev/null +++ b/releasenotes/notes/aws-lambda-microvm-identity-refresh-3a672cd6bcbad16d.yaml @@ -0,0 +1,5 @@ +--- +features: + - | + internal: Rotates the runtime ID and refreshes registered components when applications + start in AWS Lambda MicroVM environments. diff --git a/tests/contrib/asgi/test_asgi.py b/tests/contrib/asgi/test_asgi.py index 26c725062d7..c2544f4c4c5 100644 --- a/tests/contrib/asgi/test_asgi.py +++ b/tests/contrib/asgi/test_asgi.py @@ -224,6 +224,42 @@ def record_span_creation(*args, **kwargs): assert events.index("request_starting") < events.index("span") +async def test_microvm_run_hook_refreshes_identity(scope): + from ddtrace.internal import _runtime_id + from ddtrace.internal import runtime + + event_name = WebFrameworkEvents.WEB_REQUEST_STARTING.value + app = TraceMiddleware(basic_app) + scope.update({"method": "POST", "path": "/run", "root_path": "/aws/lambda-microvms/runtime/v1"}) + _runtime_id._IDENTITY_REFRESH_HOOK_REFRESHED.clear() + _runtime_id._IDENTITY_REFRESH_HOOK_RUNTIME_ID = None + core.reset_listeners(event_name, runtime.maybe_refresh_identity) + + try: + with mock.patch.object(_runtime_id, "in_aws_lambda_microvm", return_value=True): + runtime.listen_for_identity_refresh_hooks(core.on) + runtime_id = runtime.get_runtime_id() + + instance = ApplicationCommunicator(app, scope) + await instance.send_input({"type": "http.request", "body": b""}) + await instance.receive_output(1) + await instance.receive_output(1) + + refreshed_runtime_id = runtime.get_runtime_id() + assert refreshed_runtime_id != runtime_id + + instance = ApplicationCommunicator(app, scope) + await instance.send_input({"type": "http.request", "body": b""}) + await instance.receive_output(1) + await instance.receive_output(1) + + assert runtime.get_runtime_id() == refreshed_runtime_id + finally: + core.reset_listeners(event_name, runtime.maybe_refresh_identity) + _runtime_id._IDENTITY_REFRESH_HOOK_REFRESHED.clear() + _runtime_id._IDENTITY_REFRESH_HOOK_RUNTIME_ID = None + + @pytest.mark.asyncio async def test_basic_asgi(scope, test_spans): app = TraceMiddleware(basic_app) diff --git a/tests/contrib/wsgi/test_wsgi.py b/tests/contrib/wsgi/test_wsgi.py index 694ba8bfac3..11f5fc7808e 100644 --- a/tests/contrib/wsgi/test_wsgi.py +++ b/tests/contrib/wsgi/test_wsgi.py @@ -152,6 +152,28 @@ def test_web_request_starting_does_not_dispatch_without_a_listener(): dispatch.assert_not_called() +def test_web_request_starting_isolates_listener_errors(): + event_name = WebFrameworkEvents.WEB_REQUEST_STARTING.value + + def fail_request_starting(request_method, request_path): + raise RuntimeError("boom") + + original_raise = config._raise + config._raise = False + core.on(event_name, fail_request_starting) + try: + trace_utils.dispatch_wsgi_web_request_starting( + { + "REQUEST_METHOD": "POST", + "SCRIPT_NAME": "/aws/lambda-microvms/runtime/v1", + "PATH_INFO": "/run", + } + ) + finally: + core.reset_listeners(event_name, fail_request_starting) + config._raise = original_raise + + def test_web_request_starting_dispatch_precedes_span_creation(tracer): events = [] original_context_with_data = wsgi_module.core.context_with_data @@ -178,6 +200,37 @@ def record_context_with_data(*args, **kwargs): assert events.index("request_starting") < events.index("span") +def test_microvm_run_hook_refreshes_identity(tracer): + from ddtrace.internal import _runtime_id + from ddtrace.internal import runtime + + event_name = WebFrameworkEvents.WEB_REQUEST_STARTING.value + app = TestApp(DDWSGIMiddleware(application, tracer=tracer)) + _runtime_id._IDENTITY_REFRESH_HOOK_REFRESHED.clear() + _runtime_id._IDENTITY_REFRESH_HOOK_RUNTIME_ID = None + core.reset_listeners(event_name, runtime.maybe_refresh_identity) + + try: + with mock.patch.object(_runtime_id, "in_aws_lambda_microvm", return_value=True): + runtime.listen_for_identity_refresh_hooks(core.on) + runtime_id = runtime.get_runtime_id() + + resp = app.post("/run", extra_environ={"SCRIPT_NAME": "/aws/lambda-microvms/runtime/v1"}) + + assert resp.status == "200 OK" + refreshed_runtime_id = runtime.get_runtime_id() + assert refreshed_runtime_id != runtime_id + + resp = app.post("/run", extra_environ={"SCRIPT_NAME": "/aws/lambda-microvms/runtime/v1"}) + + assert resp.status == "200 OK" + assert runtime.get_runtime_id() == refreshed_runtime_id + finally: + core.reset_listeners(event_name, runtime.maybe_refresh_identity) + _runtime_id._IDENTITY_REFRESH_HOOK_REFRESHED.clear() + _runtime_id._IDENTITY_REFRESH_HOOK_RUNTIME_ID = None + + def test_middleware(tracer, test_spans): app = TestApp(DDWSGIMiddleware(application, tracer=tracer)) resp = app.get("/") diff --git a/tests/tracer/runtime/test_runtime_id.py b/tests/tracer/runtime/test_runtime_id.py index 4254cc4fa97..1b1c0ca4e20 100644 --- a/tests/tracer/runtime/test_runtime_id.py +++ b/tests/tracer/runtime/test_runtime_id.py @@ -1,3 +1,5 @@ +import os + import pytest @@ -11,6 +13,14 @@ def test_get_runtime_id(): assert runtime_id == runtime.get_runtime_id() +def test_runtime_identity_refresh_lock_is_reexported(): + from ddtrace.internal import runtime + + assert "get_runtime_identity_refresh_lock" in runtime.__all__ + with runtime.get_runtime_identity_refresh_lock(): + pass + + @pytest.mark.subprocess(env={"PYTHONWARNINGS": "ignore::DeprecationWarning"}) def test_get_runtime_id_fork(): import os @@ -433,6 +443,473 @@ def on_change(self, new_id): assert status == 0, err +@pytest.mark.parametrize("auto_enable_crashtracking", [False]) +def test_listen_for_identity_refresh_hooks_noop_does_not_import_core(monkeypatch, auto_enable_crashtracking): + import builtins + + import ddtrace.internal._runtime_id as runtime_impl + import ddtrace.internal.runtime as runtime + + monkeypatch.setattr(runtime_impl, "in_aws_lambda_microvm", lambda: False) + + real_import = builtins.__import__ + + def fail_core_import(name, *args, **kwargs): + fromlist = kwargs.get("fromlist", ()) + if len(args) >= 3: + fromlist = args[2] + if name == "ddtrace.internal" and "core" in fromlist: + raise AssertionError("listen_for_identity_refresh_hooks() imported core outside a MicroVM") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", fail_core_import) + + runtime.listen_for_identity_refresh_hooks(lambda _event, _callback: None) + + +def test_web_request_starting_event_name_is_shared_with_contrib_event(): + from ddtrace.contrib._events.web_framework import WebFrameworkEvents + from ddtrace.internal.constants import WEB_REQUEST_STARTING_EVENT + + assert WEB_REQUEST_STARTING_EVENT == WebFrameworkEvents.WEB_REQUEST_STARTING.value + + +@pytest.mark.subprocess(env={"AWS_LAMBDA_MICROVM_IMAGE_ARN": "arn:aws:lambda:us-east-1::runtime:python3.12"}, err=None) +def test_import_ddtrace_in_microvm_environment(): + """MicroVM hook registration must not import contrib before ddtrace.config exists.""" + import ddtrace + + assert ddtrace.config is not None + + +@pytest.mark.subprocess( + env={ + "AWS_LAMBDA_MICROVM_IMAGE_ARN": None, + "_DD_GLOBAL_TRACER_INIT": "false", + "DD_INSTRUMENTATION_TELEMETRY_ENABLED": "false", + }, + err=None, +) +def test_import_ddtrace_outside_microvm_does_not_import_core(): + """Normal ddtrace imports do not load the event core needed only by MicroVM hooks.""" + import sys + + import ddtrace # noqa: F401 + + assert "ddtrace.internal.core" not in sys.modules + + +@pytest.mark.subprocess(env={"AWS_LAMBDA_MICROVM_IMAGE_ARN": "arn:aws:lambda:us-east-1::runtime:python3.12"}, err=None) +def test_maybe_refresh_identity_matches_microvm_run_hook(): + """Only the exact AWS Lambda MicroVM /run hook request triggers a refresh.""" + import ddtrace.internal.runtime as runtime + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_METHOD + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_PATH + + runtime_id = runtime.get_runtime_id() + + runtime.maybe_refresh_identity(MICROVM_RUN_HOOK_METHOD, MICROVM_RUN_HOOK_PATH) + + refreshed_runtime_id = runtime.get_runtime_id() + assert refreshed_runtime_id != runtime_id + + runtime.maybe_refresh_identity(MICROVM_RUN_HOOK_METHOD, MICROVM_RUN_HOOK_PATH) + + assert runtime.get_runtime_id() == refreshed_runtime_id + + +@pytest.mark.subprocess(env={"AWS_LAMBDA_MICROVM_IMAGE_ARN": "arn:aws:lambda:us-east-1::runtime:python3.12"}, err=None) +def test_identity_refresh_hook_runs_before_root_span_creation(): + """The pre-request hook must refresh runtime-id before a web root span reads it.""" + from ddtrace import tracer + from ddtrace.contrib._events.web_framework import WebFrameworkEvents + from ddtrace.internal import core + import ddtrace.internal.runtime as runtime + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_METHOD + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_PATH + + runtime_id = runtime.get_runtime_id() + core.dispatch(WebFrameworkEvents.WEB_REQUEST_STARTING.value, (MICROVM_RUN_HOOK_METHOD, MICROVM_RUN_HOOK_PATH)) + + refreshed_runtime_id = runtime.get_runtime_id() + assert refreshed_runtime_id != runtime_id + + with tracer.trace("web.request") as span: + assert span.get_tag("runtime-id") == refreshed_runtime_id + + core.dispatch(WebFrameworkEvents.WEB_REQUEST_STARTING.value, (MICROVM_RUN_HOOK_METHOD, MICROVM_RUN_HOOK_PATH)) + + assert runtime.get_runtime_id() == refreshed_runtime_id + + +@pytest.mark.subprocess(env={"AWS_LAMBDA_MICROVM_IMAGE_ARN": "arn:aws:lambda:us-east-1::runtime:python3.12"}, err=None) +def test_maybe_refresh_identity_is_thread_safe(): + """Concurrent observations of the same /run hook refresh identity once.""" + import threading + import time + + import ddtrace.internal._runtime_id as runtime_impl + import ddtrace.internal.runtime as runtime + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_METHOD + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_PATH + + calls = [] + + def refresh_identity(*, raise_on_error=False): + calls.append(1) + time.sleep(0.01) + + runtime_impl._refresh_runtime_id = refresh_identity + + workers = 16 + barrier = threading.Barrier(workers) + errors = [] + threads = [] + + def refresh_from_request_layer(): + try: + barrier.wait() + runtime.maybe_refresh_identity(MICROVM_RUN_HOOK_METHOD, MICROVM_RUN_HOOK_PATH) + except Exception as e: + errors.append(e) + + for _ in range(workers): + thread = threading.Thread(target=refresh_from_request_layer) + thread.start() + threads.append(thread) + + for thread in threads: + thread.join() + + assert errors == [] + assert len(calls) == 1 + + +@pytest.mark.subprocess(env={"AWS_LAMBDA_MICROVM_IMAGE_ARN": "arn:aws:lambda:us-east-1::runtime:python3.12"}, err=None) +def test_maybe_refresh_identity_retries_failed_callback_without_rotating_runtime_id(): + """A failed callback retries without rerunning callbacks that already completed.""" + import threading + + import pytest + + from ddtrace.internal import _runtime_id as runtime_impl + import ddtrace.internal.runtime as runtime + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_METHOD + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_PATH + + successful_calls = [] + failed_calls = [] + retry_lock_attempted = threading.Event() + retry_lock_acquired = threading.Event() + retry_lock_threads = [] + + def on_refresh_successfully(new_id): + successful_calls.append(new_id) + + def on_refresh_failing_once(new_id): + failed_calls.append(new_id) + if len(failed_calls) == 1: + raise RuntimeError("identity refresh callback failed") + + def acquire_refresh_lock(): + retry_lock_attempted.set() + with runtime.get_runtime_identity_refresh_lock(): + retry_lock_acquired.set() + + retry_lock_thread = threading.Thread(target=acquire_refresh_lock) + retry_lock_thread.start() + retry_lock_threads.append(retry_lock_thread) + assert retry_lock_attempted.wait(timeout=1) + assert not retry_lock_acquired.wait(timeout=1) + + runtime.on_runtime_identity_refresh(on_refresh_successfully) + runtime.on_runtime_identity_refresh(on_refresh_failing_once) + # The production registry is a set, so use an explicit snapshot order for this test. + runtime_impl._ON_RUNTIME_IDENTITY_REFRESH = [on_refresh_successfully, on_refresh_failing_once] + + with pytest.raises(RuntimeError, match="identity refresh callback failed"): + runtime.maybe_refresh_identity(MICROVM_RUN_HOOK_METHOD, MICROVM_RUN_HOOK_PATH) + + refreshed_runtime_id = runtime.get_runtime_id() + assert successful_calls == [refreshed_runtime_id] + assert failed_calls == [refreshed_runtime_id] + + runtime.maybe_refresh_identity(MICROVM_RUN_HOOK_METHOD, MICROVM_RUN_HOOK_PATH) + + for retry_lock_thread in retry_lock_threads: + retry_lock_thread.join(timeout=1) + assert not retry_lock_thread.is_alive() + + assert successful_calls == [refreshed_runtime_id] + assert failed_calls == [refreshed_runtime_id, refreshed_runtime_id] + assert retry_lock_acquired.is_set() + assert runtime_impl._IDENTITY_REFRESH_HOOK_REFRESHED.is_set() + assert runtime_impl._IDENTITY_REFRESH_HOOK_PENDING_CALLBACKS is None + + +@pytest.mark.subprocess(env={"AWS_LAMBDA_MICROVM_IMAGE_ARN": "arn:aws:lambda:us-east-1::runtime:python3.12"}, err=None) +def test_maybe_refresh_identity_persistent_failure_does_not_starve_later_callbacks(): + """A callback that keeps failing must not stop the callbacks queued behind it from running.""" + import pytest + + from ddtrace.internal import _runtime_id as runtime_impl + import ddtrace.internal.runtime as runtime + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_METHOD + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_PATH + + failing_calls = [] + later_calls = [] + fail = True + + def on_refresh_failing(new_id): + failing_calls.append(new_id) + if fail: + raise RuntimeError("identity refresh callback failed") + + def on_refresh_later(new_id): + later_calls.append(new_id) + + runtime.on_runtime_identity_refresh(on_refresh_failing) + runtime.on_runtime_identity_refresh(on_refresh_later) + # The production registry is a set, so use an explicit snapshot order for this test. + runtime_impl._ON_RUNTIME_IDENTITY_REFRESH = [on_refresh_failing, on_refresh_later] + + with pytest.raises(RuntimeError, match="identity refresh callback failed"): + runtime.maybe_refresh_identity(MICROVM_RUN_HOOK_METHOD, MICROVM_RUN_HOOK_PATH) + + refreshed_runtime_id = runtime.get_runtime_id() + assert failing_calls == [refreshed_runtime_id] + assert later_calls == [refreshed_runtime_id] + assert not runtime_impl._IDENTITY_REFRESH_HOOK_REFRESHED.is_set() + assert runtime_impl._IDENTITY_REFRESH_HOOK_PENDING_CALLBACKS == [on_refresh_failing] + + # The failed callback is retried alone, against the same runtime ID. + with pytest.raises(RuntimeError, match="identity refresh callback failed"): + runtime.maybe_refresh_identity(MICROVM_RUN_HOOK_METHOD, MICROVM_RUN_HOOK_PATH) + + assert runtime.get_runtime_id() == refreshed_runtime_id + assert failing_calls == [refreshed_runtime_id, refreshed_runtime_id] + assert later_calls == [refreshed_runtime_id] + + fail = False + runtime.maybe_refresh_identity(MICROVM_RUN_HOOK_METHOD, MICROVM_RUN_HOOK_PATH) + + assert failing_calls == [refreshed_runtime_id] * 3 + assert later_calls == [refreshed_runtime_id] + assert runtime_impl._IDENTITY_REFRESH_HOOK_REFRESHED.is_set() + assert runtime_impl._IDENTITY_REFRESH_HOOK_PENDING_CALLBACKS is None + + +@pytest.mark.subprocess(env={"AWS_LAMBDA_MICROVM_IMAGE_ARN": "arn:aws:lambda:us-east-1::runtime:python3.12"}, err=None) +def test_maybe_refresh_identity_multiple_failures_raise_first_and_stay_pending(): + """Every failed callback stays pending, the first error propagates, and successes complete.""" + import pytest + + from ddtrace.internal import _runtime_id as runtime_impl + import ddtrace.internal.runtime as runtime + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_METHOD + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_PATH + + calls = [] + + def on_refresh_first_failure(new_id): + calls.append("first_failure") + raise RuntimeError("first failure") + + def on_refresh_ok(new_id): + calls.append("ok") + + def on_refresh_second_failure(new_id): + calls.append("second_failure") + raise ValueError("second failure") + + runtime.on_runtime_identity_refresh(on_refresh_first_failure) + runtime.on_runtime_identity_refresh(on_refresh_ok) + runtime.on_runtime_identity_refresh(on_refresh_second_failure) + runtime_impl._ON_RUNTIME_IDENTITY_REFRESH = [on_refresh_first_failure, on_refresh_ok, on_refresh_second_failure] + + with pytest.raises(RuntimeError, match="first failure"): + runtime.maybe_refresh_identity(MICROVM_RUN_HOOK_METHOD, MICROVM_RUN_HOOK_PATH) + + assert calls == ["first_failure", "ok", "second_failure"] + assert runtime_impl._IDENTITY_REFRESH_HOOK_PENDING_CALLBACKS == [ + on_refresh_first_failure, + on_refresh_second_failure, + ] + assert not runtime_impl._IDENTITY_REFRESH_HOOK_REFRESHED.is_set() + + +@pytest.mark.subprocess(env={"AWS_LAMBDA_MICROVM_IMAGE_ARN": None}, err=None) +def test_refresh_identity_raise_on_error_without_pending_queue_stops_at_first_failure(): + """Only the MicroVM pending queue defers errors; a direct raising refresh keeps failing fast.""" + import pytest + + from ddtrace.internal import _runtime_id as runtime_impl + import ddtrace.internal.runtime as runtime + + calls = [] + + def on_refresh_failing(new_id): + calls.append("failing") + raise RuntimeError("refresh callback failed") + + def on_refresh_later(new_id): + calls.append("later") + + runtime.on_runtime_identity_refresh(on_refresh_failing) + runtime.on_runtime_identity_refresh(on_refresh_later) + runtime_impl._ON_RUNTIME_IDENTITY_REFRESH = [on_refresh_failing, on_refresh_later] + + with pytest.raises(RuntimeError, match="refresh callback failed"): + runtime.refresh_identity(raise_on_error=True) + + assert calls == ["failing"] + + +@pytest.mark.subprocess(env={"AWS_LAMBDA_MICROVM_IMAGE_ARN": "arn:aws:lambda:us-east-1::runtime:python3.12"}, err=None) +def test_identity_refresh_failure_propagates_from_web_request_starting_dispatch(): + """A MicroVM callback failure reaches the request-start dispatch path.""" + import pytest + + from ddtrace.contrib.internal import trace_utils + from ddtrace.internal import _runtime_id as runtime_impl + import ddtrace.internal.runtime as runtime + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_PATH + + def on_refresh(_new_id): + raise RuntimeError("identity refresh callback failed") + + runtime.on_runtime_identity_refresh(on_refresh) + + with pytest.raises(RuntimeError, match="identity refresh callback failed"): + trace_utils.dispatch_wsgi_web_request_starting( + {"REQUEST_METHOD": "POST", "SCRIPT_NAME": "", "PATH_INFO": MICROVM_RUN_HOOK_PATH} + ) + + assert not runtime_impl._IDENTITY_REFRESH_HOOK_REFRESHED.is_set() + + +@pytest.mark.subprocess( + env={ + "AWS_LAMBDA_MICROVM_IMAGE_ARN": "arn:aws:lambda:us-east-1::runtime:python3.12", + "DD_TESTING_RAISE": "false", + }, + err=lambda s: "failed and stays pending" in s, +) +def test_identity_refresh_failure_is_logged_when_dispatch_suppresses_it(): + """Request-start dispatch swallows listener errors in production, so the deferred failure must be logged.""" + from ddtrace.contrib.internal import trace_utils + from ddtrace.internal import _runtime_id as runtime_impl + import ddtrace.internal.runtime as runtime + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_PATH + + def on_refresh(_new_id): + raise RuntimeError("identity refresh callback failed") + + runtime.on_runtime_identity_refresh(on_refresh) + + trace_utils.dispatch_wsgi_web_request_starting( + {"REQUEST_METHOD": "POST", "SCRIPT_NAME": "", "PATH_INFO": MICROVM_RUN_HOOK_PATH} + ) + + assert not runtime_impl._IDENTITY_REFRESH_HOOK_REFRESHED.is_set() + + +@pytest.mark.subprocess(env={"AWS_LAMBDA_MICROVM_IMAGE_ARN": "arn:aws:lambda:us-east-1::runtime:python3.12"}, err=None) +def test_microvm_refresh_guard_resets_after_fork(): + """A child process must not inherit a locked or already-refreshed MicroVM guard.""" + from ddtrace.internal import forksafe + import ddtrace.internal._runtime_id as runtime_impl + + runtime_impl._IDENTITY_REFRESH_HOOK_REFRESH_LOCK.acquire() + runtime_impl._IDENTITY_REFRESH_HOOK_REFRESHED.set() + + forksafe.ddtrace_after_in_child() + + assert runtime_impl._IDENTITY_REFRESH_HOOK_REFRESH_LOCK.acquire(False) + assert not runtime_impl._IDENTITY_REFRESH_HOOK_REFRESHED.is_set() + + +@pytest.mark.subprocess(env={"AWS_LAMBDA_MICROVM_IMAGE_ARN": "arn:aws:lambda:us-east-1::runtime:python3.12"}, err=None) +def test_maybe_refresh_identity_ignores_other_requests(): + """A different method/path, or the /resume hook, must not trigger a refresh.""" + import ddtrace.internal.runtime as runtime + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_METHOD + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_PATH + + runtime_id = runtime.get_runtime_id() + + runtime.maybe_refresh_identity("GET", MICROVM_RUN_HOOK_PATH) + runtime.maybe_refresh_identity(MICROVM_RUN_HOOK_METHOD, "/aws/lambda-microvms/runtime/v1/resume") + runtime.maybe_refresh_identity(MICROVM_RUN_HOOK_METHOD, "/some/other/path") + runtime.maybe_refresh_identity(None, MICROVM_RUN_HOOK_PATH) + runtime.maybe_refresh_identity(MICROVM_RUN_HOOK_METHOD, None) + + assert runtime.get_runtime_id() == runtime_id + + +@pytest.mark.subprocess(env={"AWS_LAMBDA_MICROVM_IMAGE_ARN": None}, err=None) +def test_maybe_refresh_identity_noop_outside_microvm(): + """Direct identity refresh calls are ignored outside a MicroVM.""" + from ddtrace.internal import _runtime_id as runtime_impl + import ddtrace.internal.runtime as runtime + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_METHOD + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_PATH + + runtime_id = runtime.get_runtime_id() + runtime.maybe_refresh_identity(MICROVM_RUN_HOOK_METHOD, MICROVM_RUN_HOOK_PATH) + + assert runtime.get_runtime_id() == runtime_id + assert runtime_impl._IDENTITY_REFRESH_HOOK_RUNTIME_ID is None + assert runtime_impl._IDENTITY_REFRESH_HOOK_PENDING_CALLBACKS is None + assert not runtime_impl._IDENTITY_REFRESH_HOOK_REFRESHED.is_set() + + +@pytest.mark.subprocess(env={"AWS_LAMBDA_MICROVM_IMAGE_ARN": None}, err=None) +def test_listen_for_identity_refresh_hooks_noop_outside_microvm(): + """Outside a MicroVM, do not register the request-event listener.""" + from ddtrace.contrib._events.web_framework import WebFrameworkEvents + from ddtrace.internal import core + import ddtrace.internal.runtime as runtime + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_METHOD + from ddtrace.internal.serverless import MICROVM_RUN_HOOK_PATH + + core.reset_listeners(WebFrameworkEvents.WEB_REQUEST_STARTING.value) + runtime.listen_for_identity_refresh_hooks(core.on) + + runtime_id = runtime.get_runtime_id() + + core.dispatch(WebFrameworkEvents.WEB_REQUEST_STARTING.value, (MICROVM_RUN_HOOK_METHOD, MICROVM_RUN_HOOK_PATH)) + + assert runtime.get_runtime_id() == runtime_id + + +@pytest.mark.parametrize("microvm_image_arn", ["", " "]) +def test_listen_for_identity_refresh_hooks_noop_for_blank_microvm_env(run_python_code_in_subprocess, microvm_image_arn): + """Blank MicroVM image ARN env values must not enable hook registration.""" + code = """ +from ddtrace.contrib._events.web_framework import WebFrameworkEvents +from ddtrace.internal import core +import ddtrace.internal.runtime as runtime +from ddtrace.internal.serverless import MICROVM_RUN_HOOK_METHOD +from ddtrace.internal.serverless import MICROVM_RUN_HOOK_PATH + +core.reset_listeners(WebFrameworkEvents.WEB_REQUEST_STARTING.value) +runtime.listen_for_identity_refresh_hooks(core.on) + +runtime_id = runtime.get_runtime_id() +core.dispatch( + WebFrameworkEvents.WEB_REQUEST_STARTING.value, (MICROVM_RUN_HOOK_METHOD, MICROVM_RUN_HOOK_PATH) +) + +assert runtime.get_runtime_id() == runtime_id +""" + env = os.environ.copy() + env["AWS_LAMBDA_MICROVM_IMAGE_ARN"] = microvm_image_arn + _, err, status, _ = run_python_code_in_subprocess(code, env=env) + assert status == 0, err + + def test_refresh_identity_notifies_refresh_subscribers(run_python_code_in_subprocess): code = """ from ddtrace.internal import runtime