Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 9 additions & 1 deletion ddtrace/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand Down
142 changes: 133 additions & 9 deletions ddtrace/internal/_runtime_id.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,10 @@
import typing as t
import uuid

from ddtrace.internal.constants import WEB_REQUEST_STARTING_EVENT
Comment thread
litianningdatadog marked this conversation as resolved.
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
Expand All @@ -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",
]

Expand All @@ -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:
Expand All @@ -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:
Expand All @@ -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:
Expand All @@ -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:
Expand All @@ -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)
Comment thread
litianningdatadog marked this conversation as resolved.
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:
Expand Down
1 change: 1 addition & 0 deletions ddtrace/internal/constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"


Expand Down
6 changes: 6 additions & 0 deletions ddtrace/internal/runtime/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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",
]
5 changes: 5 additions & 0 deletions ddtrace/internal/serverless/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
---
features:
- |
internal: Rotates the runtime ID and refreshes registered components when applications
start in AWS Lambda MicroVM environments.
36 changes: 36 additions & 0 deletions tests/contrib/asgi/test_asgi.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
53 changes: 53 additions & 0 deletions tests/contrib/wsgi/test_wsgi.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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("/")
Expand Down
Loading
Loading