diff --git a/sentry_sdk/integrations/aiohttp.py b/sentry_sdk/integrations/aiohttp.py index 858bf273f2..a0ccd9c59b 100644 --- a/sentry_sdk/integrations/aiohttp.py +++ b/sentry_sdk/integrations/aiohttp.py @@ -49,6 +49,7 @@ capture_internal_exceptions, ensure_integration_enabled, event_from_exception, + get_aws_sigv4_signed_headers, has_data_collection_enabled, logger, parse_url, @@ -456,12 +457,19 @@ async def on_request_start( span = legacy_span if should_propagate_trace(client, str(params.url)): + signed_headers = get_aws_sigv4_signed_headers( + params.headers, str(params.url) + ) for ( key, value, ) in sentry_sdk.get_current_scope().iter_trace_propagation_headers( span=span ): + # do not modify a header whose value is already signed. + if key.lower() in signed_headers: + continue + logger.debug( "[Tracing] Adding `{key}` header {value} to outgoing request to {url}.".format( key=key, value=value, url=params.url diff --git a/sentry_sdk/integrations/boto3.py b/sentry_sdk/integrations/boto3.py index 69deefc7b7..844ec5f518 100644 --- a/sentry_sdk/integrations/boto3.py +++ b/sentry_sdk/integrations/boto3.py @@ -6,8 +6,12 @@ from sentry_sdk.integrations import DidNotEnable, Integration, _check_minimum_version from sentry_sdk.scope import should_send_default_pii from sentry_sdk.traces import StreamedSpan -from sentry_sdk.tracing import Span -from sentry_sdk.tracing_utils import has_span_streaming_enabled +from sentry_sdk.tracing import BAGGAGE_HEADER_NAME, Span +from sentry_sdk.tracing_utils import ( + add_sentry_baggage_to_headers, + has_span_streaming_enabled, + should_propagate_trace, +) from sentry_sdk.utils import ( capture_internal_exceptions, parse_url, @@ -49,6 +53,8 @@ def sentry_patched_init( "request-created", partial(_sentry_request_created, service_id=service_id), ) + # run after other `before-sign` handlers, allowing it to see and preserve existing baggage. + meta.events.register_last("before-sign", _sentry_before_sign) meta.events.register("after-call", _sentry_after_call) meta.events.register("after-call-error", _sentry_after_call_error) @@ -114,6 +120,53 @@ def _sentry_request_created( request.context["_sentrysdk_span"] = span +def _sentry_before_sign( + request: "AWSRequest", signature_version: "Any", **kwargs: "Any" +) -> None: + client = sentry_sdk.get_client() + if client.get_integration(Boto3Integration) is None: + return + + with capture_internal_exceptions(): + # presigned requests are executed later by another caller. Adding propagation + # headers here would make those headers part of the signature, requiring the caller to reproduce the same values. + if isinstance(signature_version, str) and signature_version.endswith( + ("-query", "-presign-post") + ): + return + + if request.url is None or not should_propagate_trace(client, request.url): + return + + def _replace_header(request: "AWSRequest", key: str, value: str) -> None: + if key in request.headers: + del request.headers[key] + request.headers[key] = value + + # use span associated with this botocore request + span = request.context.get("_sentrysdk_span") + + headers = sentry_sdk.get_current_scope().iter_trace_propagation_headers( + span=span + ) + for header_name, header_value in headers: + if header_name != BAGGAGE_HEADER_NAME: + # normal headers (e.g. `sentry-trace`) are non-shared, so replace stale values + _replace_header(request, header_name, header_value) + continue + + # merge existing `baggage` values under single header + existing_values = request.headers.get_all(BAGGAGE_HEADER_NAME, []) + combined_baggage = { + BAGGAGE_HEADER_NAME: ",".join(str(value) for value in existing_values) + } + # preserve third-party baggage, replace stale `sentry-*` values + add_sentry_baggage_to_headers(combined_baggage, header_value) + _replace_header( + request, BAGGAGE_HEADER_NAME, combined_baggage[BAGGAGE_HEADER_NAME] + ) + + def _sentry_after_call( context: "Dict[str, Any]", parsed: "Dict[str, Any]", **kwargs: "Any" ) -> None: diff --git a/sentry_sdk/integrations/stdlib.py b/sentry_sdk/integrations/stdlib.py index 4de3819a77..3f5fff7a01 100644 --- a/sentry_sdk/integrations/stdlib.py +++ b/sentry_sdk/integrations/stdlib.py @@ -10,7 +10,7 @@ from sentry_sdk.integrations import Integration from sentry_sdk.scope import add_global_event_processor, should_send_default_pii from sentry_sdk.traces import StreamedSpan -from sentry_sdk.tracing import Span +from sentry_sdk.tracing import BAGGAGE_HEADER_NAME, Span from sentry_sdk.tracing_utils import ( EnvironHeaders, add_http_request_source, @@ -21,6 +21,7 @@ SENSITIVE_DATA_SUBSTITUTE, capture_internal_exceptions, ensure_integration_enabled, + get_aws_sigv4_signed_headers, is_sentry_url, logger, parse_url, @@ -28,7 +29,7 @@ ) if TYPE_CHECKING: - from typing import Any, Callable, Dict, List, Optional, Union + from typing import Any, Callable, Dict, List, Optional, Set, Union from sentry_sdk._types import Event, Hint @@ -61,6 +62,17 @@ def add_python_runtime_context( return event +def _request_header_names(buffer: "Optional[List[bytes]]") -> "Set[str]": + if buffer is None: + return set() + names = set() + for line in buffer: + name, separator, _ = line.partition(b":") + if separator: + names.add(name.decode("ascii", "ignore").lower()) + return names + + def _complete_span(span: "Union[Span, StreamedSpan]") -> None: if isinstance(span, StreamedSpan): with capture_internal_exceptions(): @@ -74,6 +86,7 @@ def _complete_span(span: "Union[Span, StreamedSpan]") -> None: def _install_httplib() -> None: real_putrequest = HTTPConnection.putrequest + real_endheaders = HTTPConnection.endheaders real_getresponse = HTTPConnection.getresponse real_read = HTTPResponse.read real_close = HTTPResponse.close @@ -157,26 +170,58 @@ def putrequest( set_on_span(SPANDATA.NETWORK_PEER_ADDRESS, self.host) set_on_span(SPANDATA.NETWORK_PEER_PORT, self.port) - rv = real_putrequest(self, method, url, *args, **kwargs) + try: + rv = real_putrequest(self, method, url, *args, **kwargs) + except BaseException: + self._sentrysdk_trace_url = None # type: ignore[attr-defined] + raise if should_propagate_trace(client, real_url): - for ( - key, - value, - ) in sentry_sdk.get_current_scope().iter_trace_propagation_headers( - span=span - ): - logger.debug( - "[Tracing] Adding `{key}` header {value} to outgoing request to {real_url}.".format( - key=key, value=value, real_url=real_url - ) - ) - self.putheader(key, value) + self._sentrysdk_trace_url = real_url # type: ignore[attr-defined] + else: + self._sentrysdk_trace_url = None # type: ignore[attr-defined] self._sentrysdk_span = span # type: ignore[attr-defined] return rv + def endheaders(self: "HTTPConnection", *args: "Any", **kwargs: "Any") -> "Any": + real_url = getattr(self, "_sentrysdk_trace_url", None) + span = getattr(self, "_sentrysdk_span", None) + + try: + if real_url is not None: + with capture_internal_exceptions(): + request_buffer = getattr(self, "_buffer", None) + existing_headers = _request_header_names(request_buffer) + signed_headers = get_aws_sigv4_signed_headers( + request_buffer, real_url + ) + + for ( + header_name, + header_value, + ) in sentry_sdk.get_current_scope().iter_trace_propagation_headers( + span=span + ): + normalized_header = header_name.lower() + # preserve signed headers and avoid duplicate `sentry-trace`. + if normalized_header in existing_headers and ( + normalized_header != BAGGAGE_HEADER_NAME + or normalized_header in signed_headers + ): + continue + + logger.debug( + "[Tracing] Adding `{key}` header {value} to outgoing request to {real_url}.".format( + key=header_name, value=header_value, real_url=real_url + ) + ) + self.putheader(header_name, header_value) + return real_endheaders(self, *args, **kwargs) + finally: + self._sentrysdk_trace_url = None # type: ignore[attr-defined] + def getresponse(self: "HTTPConnection", *args: "Any", **kwargs: "Any") -> "Any": span = getattr(self, "_sentrysdk_span", None) @@ -233,6 +278,7 @@ def close(self: "HTTPResponse") -> None: _complete_span(span) HTTPConnection.putrequest = putrequest # type: ignore[method-assign] + HTTPConnection.endheaders = endheaders # type: ignore[method-assign] HTTPConnection.getresponse = getresponse # type: ignore[method-assign] HTTPResponse.read = read # type: ignore[method-assign] HTTPResponse.close = close # type: ignore[assignment,method-assign] diff --git a/sentry_sdk/utils.py b/sentry_sdk/utils.py index a6ece4faf1..fb21ddfe86 100644 --- a/sentry_sdk/utils.py +++ b/sentry_sdk/utils.py @@ -1697,6 +1697,60 @@ def parse_url(url: str, sanitize: bool = True) -> "ParsedUrl": ) +def get_aws_sigv4_signed_headers( + headers: "Any", url: "Optional[str]" = None +) -> "Set[str]": + # httpConnection exposes buffer, aiohttp uses header mapping. + if isinstance(headers, (str, bytes)): + authorization = headers + elif headers is None: + authorization = "" + elif hasattr(headers, "get"): + authorization = headers.get("Authorization", "") + else: + authorization = "" + for line in headers: + name, separator, value = line.partition(b":") + if separator and name.lower() == b"authorization": + authorization = value + break + + if isinstance(authorization, bytes): + authorization = authorization.decode("ascii", "ignore") + + signed_headers: "Set[str]" = set() + if isinstance(authorization, str): + # only AWS SigV4 authorization has the SignedHeaders parameter. + value = authorization.lstrip() + if value.startswith(("AWS4-HMAC-SHA256", "AWS4-ECDSA-P256-SHA256")): + for part in value.split(","): + part = part.strip() + if part.startswith("SignedHeaders="): + _, _, header_names = part.partition("=") + signed_headers.update( + header.lower() for header in header_names.split(";") if header + ) + break + + if url is None: + return signed_headers + + query = { + key.lower(): values for key, values in parse_qs(urlsplit(url).query).items() + } + algorithm = query.get("x-amz-algorithm", [""])[0] + if algorithm not in ("AWS4-HMAC-SHA256", "AWS4-ECDSA-P256-SHA256"): + return signed_headers + + # presigned requests have SignedHeaders in the URL query. + signed_headers.update( + header.lower() + for header in query.get("x-amz-signedheaders", [""])[0].split(";") + if header + ) + return signed_headers + + def is_valid_sample_rate(rate: "Any", source: str) -> bool: """ Checks the given sample rate to make sure it is valid type and value (a diff --git a/tests/integrations/aiohttp/test_aiohttp.py b/tests/integrations/aiohttp/test_aiohttp.py index f70964e6dd..01519ed1cd 100644 --- a/tests/integrations/aiohttp/test_aiohttp.py +++ b/tests/integrations/aiohttp/test_aiohttp.py @@ -24,7 +24,7 @@ AioHttpIntegration, create_trace_config, ) -from sentry_sdk.utils import SENSITIVE_DATA_SUBSTITUTE +from sentry_sdk.utils import SENSITIVE_DATA_SUBSTITUTE, get_aws_sigv4_signed_headers from tests.conftest import ApproxDict from tests.integrations.utils import DATA_COLLECTION_USER_INFO_CASES @@ -640,6 +640,194 @@ async def handler(request): ) +@pytest.mark.asyncio +@pytest.mark.parametrize("span_streaming", [False, True]) +async def test_outgoing_trace_headers_preserve_signed_headers( + sentry_init, aiohttp_raw_server, aiohttp_client, span_streaming +): + sentry_init( + integrations=[AioHttpIntegration()], + traces_sample_rate=1.0, + trace_lifecycle="stream" if span_streaming else "static", + ) + + received_headers = [] + + async def handler(request): + received_headers.append(request.headers) + return web.Response(text="OK") + + raw_server = await aiohttp_raw_server(handler) + authorization = ( + "AWS4-HMAC-SHA256 " + "Credential=test/20260804/eu-west-1/secretsmanager/aws4_request, " + "SignedHeaders=baggage;host;sentry-trace, " + "Signature=sixtyseven" + ) + + client = await aiohttp_client(raw_server) + + if span_streaming: + with sentry_sdk.traces.start_span(name="test"): # type: ignore[attr-defined] + response = await client.get( + "/", + headers={ + "Authorization": authorization, + "baggage": "vendor=value", + "sentry-trace": "existing-trace", + }, + ) + else: + with sentry_sdk.start_transaction(name="test", sampled=True): + response = await client.get( + "/", + headers={ + "Authorization": authorization, + "baggage": "vendor=value", + "sentry-trace": "existing-trace", + }, + ) + + request_headers = received_headers[0] + + assert response.status == 200 + # signed `baggage` and `sentry-trace` headers are preserved. + assert len(request_headers.getall("baggage")) == 1 + assert request_headers["baggage"] == "vendor=value" + assert len(request_headers.getall("sentry-trace")) == 1 + assert request_headers["sentry-trace"] == "existing-trace" + assert get_aws_sigv4_signed_headers(headers=request_headers) >= { + "host", + "baggage", + "sentry-trace", + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("span_streaming", [False, True]) +async def test_outgoing_trace_headers_add_unsigned_headers( + sentry_init, aiohttp_raw_server, aiohttp_client, span_streaming +): + sentry_init( + integrations=[AioHttpIntegration()], + traces_sample_rate=1.0, + trace_lifecycle="stream" if span_streaming else "static", + ) + + received_headers = [] + + async def handler(request): + received_headers.append(request.headers.copy()) + return web.Response(text="OK") + + raw_server = await aiohttp_raw_server(handler) + authorization = ( + "AWS4-HMAC-SHA256 " + "Credential=test/20260804/eu-west-1/secretsmanager/aws4_request, " + "SignedHeaders=host;x-amz-date," + "Signature=sixtyseven" + ) + + client = await aiohttp_client(raw_server) + + if span_streaming: + with sentry_sdk.traces.start_span(name="test"): # type: ignore[attr-defined] + response = await client.get( + "/", + headers={ + "Authorization": authorization, + "baggage": "vendor=value", + "sentry-trace": "existing-trace", + }, + ) + else: + with sentry_sdk.start_transaction(name="test", sampled=True): + response = await client.get( + "/", + headers={ + "Authorization": authorization, + "baggage": "vendor=value", + "sentry-trace": "existing-trace", + }, + ) + request_headers = received_headers[0] + + assert response.status == 200 + # unsigned `baggage` and `sentry-trace` headers are injected. + assert len(request_headers.getall("sentry-trace")) == 1 + assert request_headers["sentry-trace"] != "existing-trace" + assert len(request_headers.getall("baggage")) == 1 + assert request_headers["baggage"].startswith("vendor=value,") + assert request_headers["baggage"].count("sentry-trace_id=") == 1 + # added `baggage` and `sentry-trace` are not in the signed header set. + assert get_aws_sigv4_signed_headers(headers=request_headers) == { + "host", + "x-amz-date", + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("span_streaming", [False, True]) +async def test_outgoing_trace_headers_skip_query_signed_baggage( + sentry_init, aiohttp_raw_server, aiohttp_client, span_streaming +): + sentry_init( + integrations=[AioHttpIntegration()], + traces_sample_rate=1.0, + trace_lifecycle="stream" if span_streaming else "static", + default_integrations=False, + ) + + received_headers = [] + received_urls = [] + + async def handler(request): + received_headers.append(request.headers.copy()) + # Use the raw target so aiohttp/yarl URL canonicalization does not + # change the encoding of the presigned query parameters. + received_urls.append(request.raw_path) + return web.Response(text="OK") + + raw_server = await aiohttp_raw_server(handler) + path = ( + "/" + "?X-Amz-Algorithm=AWS4-HMAC-SHA256" + "&X-Amz-Credential=" + "test%2F20260804%2Feu-west-1%2Fs3%2Faws4_request" + "&X-Amz-Date=20260804T120000Z" + "&X-Amz-Expires=60" + "&X-Amz-SignedHeaders=baggage%3Bhost" + "&X-Amz-Signature=sixtyseven" + ) + + client = await aiohttp_client(raw_server) + + if span_streaming: + with sentry_sdk.traces.start_span(name="test"): # type: ignore[attr-defined] + response = await client.get( + path, + headers={"baggage": "vendor=value"}, + ) + else: + with sentry_sdk.start_transaction(name="test", sampled=True): + response = await client.get( + path, + headers={"baggage": "vendor=value"}, + ) + + headers = received_headers[0] + assert response.status == 200 + # `baggage` is part of X-Amz-SignedHeaders, so it must not be modified. + assert len(headers.getall("baggage")) == 1 + assert headers["baggage"] == "vendor=value" + # `sentry-trace` was not signed, so it can be propagated. + assert len(headers.getall("sentry-trace")) == 1 + assert get_aws_sigv4_signed_headers(headers=headers, url=received_urls[0]) == { + "baggage", + "host", + } + + @pytest.mark.asyncio async def test_request_source_disabled( sentry_init, diff --git a/tests/integrations/boto3/test_trace_propagation.py b/tests/integrations/boto3/test_trace_propagation.py new file mode 100644 index 0000000000..c97dcf5ea4 --- /dev/null +++ b/tests/integrations/boto3/test_trace_propagation.py @@ -0,0 +1,207 @@ +from http.server import BaseHTTPRequestHandler, HTTPServer +from threading import Thread +from urllib.parse import parse_qs, urlparse + +import boto3 +import pytest +from botocore.config import Config + +import sentry_sdk +from sentry_sdk.integrations.boto3 import Boto3Integration +from sentry_sdk.integrations.stdlib import StdlibIntegration +from sentry_sdk.utils import get_aws_sigv4_signed_headers + + +class _AwsRequestHandler(BaseHTTPRequestHandler): + requests = [] + + def do_HEAD(self): + self.__class__.requests.append(self.headers) + self.send_response(200) + self.end_headers() + + def log_message(self, format, *args): + pass + + +def _start_server(): + _AwsRequestHandler.requests = [] + server = HTTPServer(("127.0.0.1", 0), _AwsRequestHandler) + thread = Thread(target=server.serve_forever, daemon=True) + thread.start() + return server, thread + + +@pytest.mark.parametrize("span_streaming", [False, True]) +def test_botocore_merges_propagation_before_sigv4_signing(sentry_init, span_streaming): + sentry_init( + traces_sample_rate=1.0, + trace_lifecycle="stream" if span_streaming else "static", + default_integrations=False, + integrations=[Boto3Integration(), StdlibIntegration()], + ) + + server, thread = _start_server() + + try: + client = boto3.client( # type: ignore[attr-defined] + "s3", + # connect to mock AWS server. + endpoint_url=f"http://127.0.0.1:{server.server_port}", + aws_access_key_id="test-access-key", + aws_secret_access_key="test-secret-key", + config=Config(signature_version="v4"), + ) + + def _inject_third_party_baggage(request, **kwargs): + request.headers.add_header( + "baggage", + "dd-origin=synthetics,sentry-trace_id=stale,sentry-sample_rand=0.100000", + ) + request.headers.add_header("baggage", "vendor=value") + + signed_request_headers = {} + + def capture_headers_after_instrumentation(request, **kwargs): + for header_name in ("baggage", "sentry-trace"): + signed_request_headers[header_name] = request.headers.get_all( + header_name + ) + + # register `before-sign` handler that adds third-party baggage. + client.meta.events.register("before-sign", _inject_third_party_baggage) + client.meta.events.register_last( + "before-sign", capture_headers_after_instrumentation + ) + + if span_streaming: + with sentry_sdk.traces.start_span( # type: ignore[attr-defined] + name="incoming" + ): + response = client.head_object( + Bucket="example-bucket", + Key="example-key", + ) + else: + with sentry_sdk.start_transaction(name="incoming", sampled=True): + response = client.head_object( + Bucket="example-bucket", + Key="example-key", + ) + + assert response["ResponseMetadata"]["HTTPStatusCode"] == 200 + headers = _AwsRequestHandler.requests[-1] + + baggage_headers = headers.get_all("baggage") + assert baggage_headers is not None + assert len(baggage_headers) == 1 + assert baggage_headers == signed_request_headers["baggage"] + + baggage = baggage_headers[0] + # preserves third-party baggage. + assert "dd-origin=synthetics" in baggage + assert "vendor=value" in baggage + # add own `sentry-*` baggage. + assert "sentry-trace_id=" in baggage + assert "sentry-trace_id=stale" not in baggage + # replace stale values instead of duplicating them. + assert baggage.count("sentry-trace_id=") == 1 + assert baggage.count("sentry-sample_rand=") == 1 + + # adds single `sentry-trace` header. + sentry_trace_headers = headers.get_all("sentry-trace") + assert sentry_trace_headers is not None + assert len(sentry_trace_headers) == 1 + assert sentry_trace_headers == signed_request_headers["sentry-trace"] + # both `baggage` and `sentry-trace` are signed. + signed_headers = get_aws_sigv4_signed_headers(headers=headers) + assert signed_headers >= {"baggage", "sentry-trace"} + finally: + server.shutdown() + server.server_close() + thread.join() + + +@pytest.mark.parametrize("span_streaming", [False, True]) +def test_botocore_without_boto3_integration_preserves_signed_baggage( + sentry_init, span_streaming +): + sentry_init( + traces_sample_rate=1.0, + trace_lifecycle="stream" if span_streaming else "static", + default_integrations=False, + integrations=[StdlibIntegration()], + ) + + server, thread = _start_server() + try: + client = boto3.client( # type: ignore[attr-defined] + "s3", + endpoint_url=f"http://127.0.0.1:{server.server_port}", + aws_access_key_id="test-access-key", + aws_secret_access_key="test-secret-key", + config=Config(signature_version="v4"), + ) + + def _inject_signed_baggage(request, **kwargs): + request.headers.add_header("baggage", "vendor=value") + + # register `before-sign` handler that third-party signed baggage. + client.meta.events.register("before-sign", _inject_signed_baggage) + + if span_streaming: + with sentry_sdk.traces.start_span( # type: ignore[attr-defined] + name="incoming" + ): + response = client.head_object( + Bucket="example-bucket", + Key="example-key", + ) + else: + with sentry_sdk.start_transaction(name="incoming", sampled=True): + response = client.head_object( + Bucket="example-bucket", + Key="example-key", + ) + + assert response["ResponseMetadata"]["HTTPStatusCode"] == 200 + headers = _AwsRequestHandler.requests[-1] + # preserves third-party signed baggage. + assert headers.get_all("baggage") == ["vendor=value"] + # `httplib` still adds single `sentry-trace` header. + assert len(headers.get_all("sentry-trace")) == 1 + signed_headers = get_aws_sigv4_signed_headers(headers=headers) + assert "baggage" in signed_headers + assert "sentry-trace" not in signed_headers + finally: + server.shutdown() + server.server_close() + thread.join() + + +def test_presigned_urls_do_not_require_sentry_headers(sentry_init): + sentry_init( + traces_sample_rate=1.0, + default_integrations=False, + integrations=[Boto3Integration(), StdlibIntegration()], + ) + client = boto3.client( # type: ignore[attr-defined] + "s3", + aws_access_key_id="test-access-key", + aws_secret_access_key="test-secret-key", + config=Config(signature_version="s3v4"), + ) + + url = client.generate_presigned_url( + "get_object", + Params={"Bucket": "example-bucket", "Key": "example-key"}, + ExpiresIn=60, + ) + query = parse_qs(urlparse(url).query) + + # only `host` header is signed. + assert query["X-Amz-SignedHeaders"] == ["host"] + assert get_aws_sigv4_signed_headers(headers={}, url=url) == {"host"} + # no `sentry-*` or baggage are added. + assert "sentry-trace" not in url + assert "baggage" not in url diff --git a/tests/integrations/stdlib/test_httplib.py b/tests/integrations/stdlib/test_httplib.py index 66d1e8db37..04eb404f75 100644 --- a/tests/integrations/stdlib/test_httplib.py +++ b/tests/integrations/stdlib/test_httplib.py @@ -16,6 +16,7 @@ from sentry_sdk import capture_message, continue_trace, start_transaction from sentry_sdk.consts import MATCH_ALL, SPANDATA from sentry_sdk.integrations.stdlib import StdlibIntegration +from sentry_sdk.utils import get_aws_sigv4_signed_headers from tests.conftest import ApproxDict, create_mock_http_server, get_free_port PORT = create_mock_http_server() @@ -76,6 +77,43 @@ def create_chunked_server(): CHUNKED_PORT = create_chunked_server() +@pytest.fixture +def local_http_server(): + requests = [] + + class TraceHeaderHandler(BaseHTTPRequestHandler): + def do_POST(self): + requests.append(self.headers) + self.send_response(200) + self.send_header("Content-Length", "0") + self.end_headers() + + server = HTTPServer(("127.0.0.1", 0), TraceHeaderHandler) + thread = Thread(target=server.serve_forever, daemon=True) + thread.start() + + try: + yield server, requests + finally: + server.shutdown() + server.server_close() + thread.join() + + +def _request(server, headers, path="/"): + connection = HTTPConnection("127.0.0.1", server.server_port) + connection.putrequest("POST", path) + + for key, value in headers: + connection.putheader(key, value) + + connection.endheaders() + + response = connection.getresponse() + response.read() + connection.close() + + def test_crumb_capture(sentry_init, capture_events): sentry_init(integrations=[StdlibIntegration()], send_default_pii=True) events = capture_events() @@ -526,6 +564,140 @@ def getresponse(self, *args, **kwargs): assert request_headers["baggage"] == expected_outgoing_baggage +@pytest.mark.parametrize("span_streaming", [False, True]) +def test_outgoing_trace_headers_append_to_unsigned_baggage( + sentry_init, local_http_server, span_streaming +): + sentry_init( + traces_sample_rate=1.0, + trace_lifecycle="stream" if span_streaming else "static", + default_integrations=False, + integrations=[StdlibIntegration()], + ) + server, requests = local_http_server + + with mock.patch("sentry_sdk.tracing_utils.Random.randrange", return_value=67): + if span_streaming: + with sentry_sdk.traces.start_span(name="test"): # type: ignore[attr-defined] + _request(server, [("baggage", "vendor=value")]) + else: + with sentry_sdk.start_transaction(name="test", sampled=True): + _request(server, [("baggage", "vendor=value")]) + + headers = requests[0] + + # preserve existing unsigned baggage + baggage_headers = headers.get_all("baggage") + assert baggage_headers is not None + assert len(baggage_headers) == 2 + assert baggage_headers[0] == "vendor=value" + assert baggage_headers[1].count("sentry-trace_id=") == 1 + assert "sentry-sample_rand=0.000067" in baggage_headers[1] + assert len(headers.get_all("sentry-trace")) == 1 + + +@pytest.mark.parametrize("span_streaming", [False, True]) +def test_outgoing_trace_headers_skip_signed_baggage( + sentry_init, local_http_server, span_streaming +): + sentry_init( + traces_sample_rate=1.0, + trace_lifecycle="stream" if span_streaming else "static", + default_integrations=False, + integrations=[StdlibIntegration()], + ) + server, requests = local_http_server + + # simulate AWS SigV4 request that is already signed. + authorization = ( + "AWS4-HMAC-SHA256 " + "Credential=test/20260804/eu-west-1/secretsmanager/aws4_request, " + "SignedHeaders=baggage;host;sentry-trace, " + "Signature=sixtyseven" + ) + + if span_streaming: + with sentry_sdk.traces.start_span(name="test"): # type: ignore[attr-defined] + _request( + server, + [ + ("baggage", "vendor=value"), + ("sentry-trace", "existing-trace"), + ("Authorization", authorization), + ], + ) + else: + with sentry_sdk.start_transaction(name="test", sampled=True): + _request( + server, + [ + ("baggage", "vendor=value"), + ("sentry-trace", "existing-trace"), + ("Authorization", authorization), + ], + ) + + headers = requests[0] + + # do not append baggage after SigV4 signs it. + assert headers.get_all("baggage") == ["vendor=value"] + # preserves existing `sentry-trace` header. + assert headers.get_all("sentry-trace") == ["existing-trace"] + assert get_aws_sigv4_signed_headers(headers=headers) >= { + "baggage", + "host", + "sentry-trace", + } + + +@pytest.mark.parametrize("span_streaming", [False, True]) +def test_outgoing_trace_headers_skip_query_signed_baggage( + sentry_init, local_http_server, span_streaming +): + sentry_init( + traces_sample_rate=1.0, + trace_lifecycle="stream" if span_streaming else "static", + default_integrations=False, + integrations=[StdlibIntegration()], + ) + server, requests = local_http_server + path = ( + "/" + "?X-Amz-Algorithm=AWS4-HMAC-SHA256" + "&X-Amz-Credential=" + "test%2F20260804%2Feu-west-1%2Fs3%2Faws4_request" + "&X-Amz-Date=20260804T120000Z" + "&X-Amz-Expires=60" + "&X-Amz-SignedHeaders=baggage%3Bhost" + "&X-Amz-Signature=sixtyseven" + ) + + if span_streaming: + with sentry_sdk.traces.start_span(name="test"): # type: ignore[attr-defined] + _request( + server, + [("baggage", "vendor=value")], + path=path, + ) + else: + with sentry_sdk.start_transaction(name="test", sampled=True): + _request( + server, + [("baggage", "vendor=value")], + path=path, + ) + + headers = requests[0] + # `baggage` is part of X-Amz-SignedHeaders, so may not be modified. + assert len(headers.get_all("baggage")) == 1 + assert headers["baggage"] == "vendor=value" + # `sentry-trace` was not signed, so it can be propagated. + assert len(headers.get_all("sentry-trace")) == 1 + assert get_aws_sigv4_signed_headers( + headers=headers, url=f"http://127.0.0.1:{server.server_port}{path}" + ) >= {"baggage", "host"} + + @pytest.mark.parametrize( "trace_propagation_targets,host,path,trace_propagated", [ diff --git a/tests/test_utils.py b/tests/test_utils.py index 64973ea5dd..716095c6f8 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -18,6 +18,7 @@ env_to_bool, exc_info_from_error, format_timestamp, + get_aws_sigv4_signed_headers, get_current_thread_meta, get_default_release, get_error_message, @@ -665,6 +666,35 @@ def test_default_release_empty_string(): assert release is None +@pytest.mark.parametrize( + "headers,url,expected", + [ + ( + { + "Authorization": ( + "AWS4-HMAC-SHA256 " + "Credential=test/20260804/eu-west-1/secretsmanager/aws4_request, " + "SignedHeaders=Host;X-Amz-Date, " + "Signature=sixtyseven" + ) + }, + None, + {"host", "x-amz-date"}, + ), + ( + {}, + "https://example.com/?" + "X-Amz-Algorithm=AWS4-HMAC-SHA256&" + "X-Amz-SignedHeaders=host%3Bx-amz-date&" + "X-Amz-Signature=sixtyseven", + {"host", "x-amz-date"}, + ), + ], +) +def test_get_aws_sigv4_signed_headers(headers, url, expected): + assert get_aws_sigv4_signed_headers(headers, url) == expected + + def test_get_default_release_sentry_release_env(monkeypatch): monkeypatch.setenv("SENTRY_RELEASE", "sentry-env-release") assert get_default_release() == "sentry-env-release"