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..91106eca74 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, @@ -28,7 +28,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 +61,41 @@ def add_python_runtime_context( return event +def _aws_sigv4_signed_headers(buffer: "Optional[List[bytes]]") -> "Set[str]": + if buffer is None: + return set() + for line in buffer: + name, separator, value = line.partition(b":") + if not separator or name.lower() != b"authorization": + continue + + value = value.lstrip() + if not value.startswith((b"AWS4-HMAC-SHA256", b"AWS4-ECDSA-P256-SHA256")): + continue + + for part in value.split(b","): + part = part.strip() + if part.startswith(b"SignedHeaders="): + _, _, header_names = part.partition(b"=") + return { + header.decode("ascii", "ignore").lower() + for header in header_names.split(b";") + if header + } + return set() + + +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 +109,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 +193,56 @@ 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: + request_buffer = getattr(self, "_buffer", None) + existing_headers = _request_header_names(request_buffer) + signed_headers = _aws_sigv4_signed_headers(request_buffer) + + 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 +299,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/tests/integrations/boto3/test_trace_propagation.py b/tests/integrations/boto3/test_trace_propagation.py new file mode 100644 index 0000000000..39477b1b42 --- /dev/null +++ b/tests/integrations/boto3/test_trace_propagation.py @@ -0,0 +1,205 @@ +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 + + +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"] + + authorization = headers["Authorization"] + signed_headers = authorization.split("SignedHeaders=", 1)[1].split(",", 1)[0] + # both `baggage` and `sentry-trace` are signed. + assert "baggage" in signed_headers.split(";") + assert "sentry-trace" in signed_headers.split(";") + 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 + 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"] + # 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..33aaa19aae 100644 --- a/tests/integrations/stdlib/test_httplib.py +++ b/tests/integrations/stdlib/test_httplib.py @@ -76,6 +76,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): + connection = HTTPConnection("127.0.0.1", server.server_port) + connection.putrequest("POST", "/") + + 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 +563,87 @@ 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=500000): + 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] + + baggage_headers = headers.get_all("baggage") + assert baggage_headers is not None + # preserve existing unsigned baggage + 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.500000" 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"] + + @pytest.mark.parametrize( "trace_propagation_targets,host,path,trace_propagated", [