diff --git a/packages/http/httpx/kiota_http/httpx_request_adapter.py b/packages/http/httpx/kiota_http/httpx_request_adapter.py index 20904ded..8fa85187 100644 --- a/packages/http/httpx/kiota_http/httpx_request_adapter.py +++ b/packages/http/httpx/kiota_http/httpx_request_adapter.py @@ -44,7 +44,7 @@ from ._version import VERSION from .kiota_client_factory import KiotaClientFactory -from .middleware import ParametersNameDecodingHandler +from .middleware import REQUEST_OPTIONS_KEY, ParametersNameDecodingHandler from .middleware.options import ParametersNameDecodingHandlerOption, ResponseHandlerOption from .observability_options import ObservabilityOptions @@ -686,18 +686,18 @@ def get_request_from_request_information( if self.observability_options.include_euii_attributes: otel_attributes.update({URL_FULL: url.geturl()}) + request_options = { + self.observability_options.get_key(): self.observability_options, + "parent_span": parent_span, + **request_info.request_options, + } request = self._http_client.build_request( method=method.value, url=request_info.url, headers=request_info.request_headers, content=request_info.content, + extensions={REQUEST_OPTIONS_KEY: request_options}, ) - request_options = { - self.observability_options.get_key(): self.observability_options, - "parent_span": parent_span, - **request_info.request_options, - } - setattr(request, "options", request_options) if content_length := request.headers.get("Content-Length", None): otel_attributes.update({"http.request.body.size": content_length}) diff --git a/packages/http/httpx/kiota_http/kiota_client_factory.py b/packages/http/httpx/kiota_http/kiota_client_factory.py index 9d70d3ff..afa75eb3 100644 --- a/packages/http/httpx/kiota_http/kiota_client_factory.py +++ b/packages/http/httpx/kiota_http/kiota_client_factory.py @@ -11,6 +11,7 @@ from .middleware import ( AsyncKiotaTransport, BaseMiddleware, + BodyInspectionHandler, HeadersInspectionHandler, MiddlewarePipeline, ParametersNameDecodingHandler, @@ -19,6 +20,7 @@ UrlReplaceHandler, ) from .middleware.options import ( + BodyInspectionHandlerOption, HeadersInspectionHandlerOption, ParametersNameDecodingHandlerOption, RedirectHandlerOption, @@ -91,6 +93,7 @@ def get_default_middleware(options: Optional[dict[str, RequestOption]]) -> list[ url_replace_handler = UrlReplaceHandler() user_agent_handler = UserAgentHandler() headers_inspection_handler = HeadersInspectionHandler() + body_inspection_handler = BodyInspectionHandler() if options: redirect_handler_options = options.get(RedirectHandlerOption.get_key()) @@ -135,11 +138,18 @@ def get_default_middleware(options: Optional[dict[str, RequestOption]]) -> list[ options=headers_inspection_handler_options ) - middleware = [ + body_inspection_handler_options = options.get(BodyInspectionHandlerOption.get_key()) + if body_inspection_handler_options and isinstance( + body_inspection_handler_options, BodyInspectionHandlerOption + ): + body_inspection_handler = BodyInspectionHandler( + options=body_inspection_handler_options + ) + + return [ redirect_handler, retry_handler, parameters_name_decoding_handler, url_replace_handler, - user_agent_handler, headers_inspection_handler + user_agent_handler, headers_inspection_handler, body_inspection_handler ] - return middleware @staticmethod def create_middleware_pipeline( diff --git a/packages/http/httpx/kiota_http/middleware/__init__.py b/packages/http/httpx/kiota_http/middleware/__init__.py index 1e9b94a7..853b9977 100644 --- a/packages/http/httpx/kiota_http/middleware/__init__.py +++ b/packages/http/httpx/kiota_http/middleware/__init__.py @@ -1,6 +1,7 @@ from .async_kiota_transport import AsyncKiotaTransport +from .body_inspection_handler import BodyInspectionHandler from .headers_inspection_handler import HeadersInspectionHandler -from .middleware import BaseMiddleware, MiddlewarePipeline +from .middleware import REQUEST_OPTIONS_KEY, BaseMiddleware, MiddlewarePipeline from .parameters_name_decoding_handler import ParametersNameDecodingHandler from .redirect_handler import RedirectHandler from .retry_handler import RetryHandler diff --git a/packages/http/httpx/kiota_http/middleware/body_inspection_handler.py b/packages/http/httpx/kiota_http/middleware/body_inspection_handler.py new file mode 100644 index 00000000..4c3f1a99 --- /dev/null +++ b/packages/http/httpx/kiota_http/middleware/body_inspection_handler.py @@ -0,0 +1,116 @@ +# ------------------------------------ +# Copyright (c) Microsoft Corporation. All Rights Reserved. +# Licensed under the MIT License. +# See License in the project root for license information. +# ------------------------------------ + +from typing import Optional + +import httpx + +from .middleware import REQUEST_OPTIONS_KEY, BaseMiddleware +from .options import BodyInspectionHandlerOption + +BODY_INSPECTION_KEY = "com.microsoft.kiota.handler.bodyInspection.enable" + + +class BodyInspectionHandler(BaseMiddleware): + """The Body Inspection Handler allows the developer to inspect the body of the + request and response. + """ + + def __init__( + self, + options: Optional[BodyInspectionHandlerOption] = None, + ): + """Create an instance of BodyInspectionHandler + + Args: + options (BodyInspectionHandlerOption, optional): Default options to apply to the + handler. A new BodyInspectionHandlerOption per handler when not provided. + """ + super().__init__() + self.options = options if options is not None else BodyInspectionHandlerOption() + + async def send( + self, request: httpx.Request, transport: httpx.AsyncBaseTransport + ) -> httpx.Response: + """To execute the current middleware + + Args: + request (httpx.Request): The prepared request object + transport (httpx.AsyncBaseTransport): The HTTP transport to use + + Returns: + httpx.Response: The response object. + """ + if request is None: + raise TypeError("request cannot be null") + + current_options = self._get_current_options(request) + span = self._create_observability_span(request, "BodyInspectionHandler_send") + try: + span.set_attribute(BODY_INSPECTION_KEY, True) + + if current_options and current_options.inspect_request_body: + content = await request.aread() + if content: + current_options.request_body = content + else: + current_options.request_body = None + + response = await super().send(request, transport) + + if current_options and current_options.inspect_response_body: + response_content: Optional[bytes] = None + # A consumed stream is inspectable only when HTTPX cached its content. + if hasattr(response, "_content"): + response_content = response.content + elif not response.is_stream_consumed and not response.is_closed: + num_bytes_downloaded = response.num_bytes_downloaded + raw_content = b"".join([chunk async for chunk in response.aiter_raw()]) + self._restore_response_stream(response, raw_content, num_bytes_downloaded) + response_content = await response.aread() + self._restore_response_stream(response, raw_content, num_bytes_downloaded) + if response_content: + current_options.response_body = response_content + else: + current_options.response_body = None + + return response + finally: + span.end() + + def _get_current_options(self, request: httpx.Request) -> BodyInspectionHandlerOption: + """Returns the options to use for the request. Overrides default options if + request options are passed. + + Args: + request (httpx.Request): The prepared request object + + Returns: + BodyInspectionHandlerOption: The options to be used. + """ + current_options = None + request_options = request.extensions.get(REQUEST_OPTIONS_KEY) + if request_options: + current_options = request_options.get(BodyInspectionHandlerOption.get_key(), None) + if not current_options: + current_options = self.options + + current_options._clear_captured_bodies() + return current_options + + @staticmethod + def _restore_response_stream( + response: httpx.Response, content: bytes, num_bytes_downloaded: int + ) -> None: + # aread() caches decoded content and a stateful decoder; discard both when rewinding. + if hasattr(response, "_content"): + del response._content + if hasattr(response, "_decoder"): + del response._decoder + response.stream = httpx.ByteStream(content) + response.is_stream_consumed = False + response.is_closed = False + response._num_bytes_downloaded = num_bytes_downloaded diff --git a/packages/http/httpx/kiota_http/middleware/headers_inspection_handler.py b/packages/http/httpx/kiota_http/middleware/headers_inspection_handler.py index 81b03f40..0782b76c 100644 --- a/packages/http/httpx/kiota_http/middleware/headers_inspection_handler.py +++ b/packages/http/httpx/kiota_http/middleware/headers_inspection_handler.py @@ -10,7 +10,7 @@ import httpx -from .middleware import BaseMiddleware +from .middleware import REQUEST_OPTIONS_KEY, BaseMiddleware from .options import HeadersInspectionHandlerOption HEADERS_INSPECTION_KEY = "com.microsoft.kiota.handler.headers_inspection.enable" @@ -71,9 +71,9 @@ def _get_current_options(self, request: httpx.Request) -> HeadersInspectionHandl HeadersInspectionHandlerOption: The options to be used. """ current_options = None - request_options = getattr(request, "options", None) + request_options = request.extensions.get(REQUEST_OPTIONS_KEY) if request_options: - current_options = request_options.get( # type:ignore + current_options = request_options.get( HeadersInspectionHandlerOption.get_key(), None ) if current_options: diff --git a/packages/http/httpx/kiota_http/middleware/middleware.py b/packages/http/httpx/kiota_http/middleware/middleware.py index 6e70e385..97e713f7 100644 --- a/packages/http/httpx/kiota_http/middleware/middleware.py +++ b/packages/http/httpx/kiota_http/middleware/middleware.py @@ -10,6 +10,8 @@ tracer = trace.get_tracer(ObservabilityOptions.get_tracer_instrumentation_name(), VERSION) +REQUEST_OPTIONS_KEY = "kiota_request_options" + class MiddlewarePipeline(): """MiddlewarePipeline, entry point of middleware @@ -56,8 +58,8 @@ def __init__(self): async def send(self, request, transport): if self.next is None: # Remove request options if there's no other middleware in the chain. - if hasattr(request, "options") and request.options: - delattr(request, 'options') + if hasattr(request, "extensions") and isinstance(request.extensions, dict): + request.extensions.pop(REQUEST_OPTIONS_KEY, None) response = await transport.handle_async_request(request) response.request = request return response @@ -68,11 +70,13 @@ def _create_observability_span(self, request, span_name: str) -> trace.Span: If no parent_span is found in the request, uses the parent_span in the object. If parent_span is None, current context will be used.""" _span = None - if options := getattr(request, "options", None): - if parent_span := options.get("parent_span", None): - self.parent_span = parent_span - _context = trace.set_span_in_context(parent_span) - _span = tracer.start_span(span_name, _context) + options = None + if hasattr(request, "extensions") and isinstance(request.extensions, dict): + options = request.extensions.get(REQUEST_OPTIONS_KEY) + if options and (parent_span := options.get("parent_span", None)): + self.parent_span = parent_span + _context = trace.set_span_in_context(parent_span) + _span = tracer.start_span(span_name, _context) if _span is None: _context = trace.set_span_in_context(self.parent_span) _span = tracer.start_span(span_name, _context) diff --git a/packages/http/httpx/kiota_http/middleware/options/__init__.py b/packages/http/httpx/kiota_http/middleware/options/__init__.py index 2d6faf50..186be925 100644 --- a/packages/http/httpx/kiota_http/middleware/options/__init__.py +++ b/packages/http/httpx/kiota_http/middleware/options/__init__.py @@ -1,3 +1,4 @@ +from .body_inspection_handler_option import BodyInspectionHandlerOption from .headers_inspection_handler_option import HeadersInspectionHandlerOption from .parameters_name_decoding_handler_option import ParametersNameDecodingHandlerOption from .redirect_handler_option import RedirectHandlerOption diff --git a/packages/http/httpx/kiota_http/middleware/options/body_inspection_handler_option.py b/packages/http/httpx/kiota_http/middleware/options/body_inspection_handler_option.py new file mode 100644 index 00000000..13ff69fb --- /dev/null +++ b/packages/http/httpx/kiota_http/middleware/options/body_inspection_handler_option.py @@ -0,0 +1,163 @@ +# ------------------------------------ +# Copyright (c) Microsoft Corporation. All Rights Reserved. +# Licensed under the MIT License. +# See License in the project root for license information. +# ------------------------------------ +from contextvars import ContextVar +from io import BytesIO +from typing import ClassVar, Optional +from weakref import WeakKeyDictionary + +from kiota_abstractions.request_option import RequestOption + + +class BodyInspectionHandlerOption(RequestOption): + """Config options for the BodyInspectionHandler. + + Captured bodies are isolated by execution context so concurrent requests can + safely share the same option instance. + + Args: + inspect_request_body (bool, optional): Whether the request body + should be inspected. Defaults to False. Note that this setting + increases memory usage as the request body is copied in memory. + inspect_response_body (bool, optional): Whether the response body + should be inspected. Defaults to False. Note that this setting + increases memory usage as the response body is copied in memory. + request_body (Optional[bytes], optional): The inspected request body bytes. + Defaults to None. + response_body (Optional[bytes], optional): The inspected response body bytes. + Defaults to None. + """ + + BODY_INSPECTION_HANDLER_OPTION_KEY: ClassVar[str] = "BodyInspectionHandlerOption" + + def __init__( + self, + inspect_request_body: bool = False, + inspect_response_body: bool = False, + request_body: Optional[bytes] = None, + response_body: Optional[bytes] = None, + ) -> None: + self.inspect_request_body = inspect_request_body + self.inspect_response_body = inspect_response_body + self._default_request_body = request_body + self._default_response_body = response_body + + @property + def request_body(self) -> Optional[bytes]: + """Gets the request body captured in the current execution context.""" + request_body, _ = self._get_captured_bodies() + return request_body + + @request_body.setter + def request_body(self, value: Optional[bytes]) -> None: + _, response_body = self._get_captured_bodies() + self._set_captured_bodies(value, response_body) + + @property + def response_body(self) -> Optional[bytes]: + """Gets the response body captured in the current execution context.""" + _, response_body = self._get_captured_bodies() + return response_body + + @response_body.setter + def response_body(self, value: Optional[bytes]) -> None: + request_body, _ = self._get_captured_bodies() + self._set_captured_bodies(request_body, value) + + def _get_captured_bodies(self) -> tuple[Optional[bytes], Optional[bytes]]: + captures = _BODY_CAPTURES.get() + if captures is not None: + captured_bodies = captures.get(self) + if captured_bodies is not None: + return captured_bodies + return self._default_request_body, self._default_response_body + + def _set_captured_bodies( + self, request_body: Optional[bytes], response_body: Optional[bytes] + ) -> None: + current_captures = _BODY_CAPTURES.get() + captures = ( + WeakKeyDictionary() + if current_captures is None else WeakKeyDictionary(current_captures) + ) + captures[self] = (request_body, response_body) + _BODY_CAPTURES.set(captures) + + def _clear_captured_bodies(self) -> None: + request_body, response_body = self._get_captured_bodies() + if request_body is not None or response_body is not None: + self._set_captured_bodies(None, None) + + @staticmethod + def get_key() -> str: + return BodyInspectionHandlerOption.BODY_INSPECTION_HANDLER_OPTION_KEY + + def get_request_body(self) -> Optional[bytes]: + """Gets the request body as bytes. + + Returns: + Optional[bytes]: The request body bytes, or None if inspection was + disabled or no body was present. + """ + return self.request_body + + def get_response_body(self) -> Optional[bytes]: + """Gets the response body as bytes. + + Returns: + Optional[bytes]: The response body bytes, or None if inspection was + disabled or no body was present. + """ + return self.response_body + + def get_request_body_stream(self) -> Optional[BytesIO]: + """Gets the request body as a seekable stream rewound to position 0. + + Callers are responsible for disposing/closing the stream. Note that this stream + is a copy of the original request body, which has impact on memory usage. + + Returns: + Optional[BytesIO]: A new BytesIO stream of the request body, or None if + inspection was disabled or no body was present. + """ + if self.request_body is not None: + stream = BytesIO(self.request_body) + stream.seek(0) + return stream + return None + + def get_response_body_stream(self) -> Optional[BytesIO]: + """Gets the response body as a seekable stream rewound to position 0. + + Callers are responsible for disposing/closing the stream. Note that this stream + is a copy of the original response body, which has impact on memory usage. + + Returns: + Optional[BytesIO]: A new BytesIO stream of the response body, or None if + inspection was disabled or no body was present. + """ + if self.response_body is not None: + stream = BytesIO(self.response_body) + stream.seek(0) + return stream + return None + + @property + def request_body_stream(self) -> Optional[BytesIO]: + """Stream property for the request body rewound to position 0.""" + return self.get_request_body_stream() + + @property + def response_body_stream(self) -> Optional[BytesIO]: + """Stream property for the response body rewound to position 0.""" + return self.get_response_body_stream() + + +_BodyCapture = tuple[Optional[bytes], Optional[bytes]] +_BodyCaptureMap = WeakKeyDictionary[BodyInspectionHandlerOption, _BodyCapture] +_OptionalBodyCaptureMap = Optional[_BodyCaptureMap] +_BODY_CAPTURES: ContextVar[_OptionalBodyCaptureMap] = ContextVar( + "kiota_body_inspection_captures", default=None +) diff --git a/packages/http/httpx/kiota_http/middleware/parameters_name_decoding_handler.py b/packages/http/httpx/kiota_http/middleware/parameters_name_decoding_handler.py index e7471df9..5b0f0eb7 100644 --- a/packages/http/httpx/kiota_http/middleware/parameters_name_decoding_handler.py +++ b/packages/http/httpx/kiota_http/middleware/parameters_name_decoding_handler.py @@ -2,7 +2,7 @@ import httpx -from .middleware import BaseMiddleware +from .middleware import REQUEST_OPTIONS_KEY, BaseMiddleware from .options import ParametersNameDecodingHandlerOption PARAMETERS_NAME_DECODING_KEY = "com.microsoft.kiota.handler.parameters_name_decoding.enable" @@ -70,9 +70,9 @@ def _get_current_options(self, request: httpx.Request) -> ParametersNameDecoding Returns: ParametersNameDecodingHandlerOption: The options to used. """ - request_options = getattr(request, "options", None) + request_options = request.extensions.get(REQUEST_OPTIONS_KEY) if request_options: - current_options = request_options.get( # type:ignore + current_options = request_options.get( ParametersNameDecodingHandlerOption.get_key(), self.options ) return current_options diff --git a/packages/http/httpx/kiota_http/middleware/redirect_handler.py b/packages/http/httpx/kiota_http/middleware/redirect_handler.py index 9be19db8..9c73a1ac 100644 --- a/packages/http/httpx/kiota_http/middleware/redirect_handler.py +++ b/packages/http/httpx/kiota_http/middleware/redirect_handler.py @@ -6,7 +6,7 @@ import httpx from .._exceptions import RedirectError -from .middleware import BaseMiddleware +from .middleware import REQUEST_OPTIONS_KEY, BaseMiddleware from .options import RedirectHandlerOption REDIRECT_ENABLE_KEY = "com.microsoft.kiota.handler.redirect.enable" @@ -64,6 +64,7 @@ async def send( """ _enable_span = self._create_observability_span(request, "RedirectHandler_send") current_options = self._get_current_options(request) + request_options = request.extensions.get(REQUEST_OPTIONS_KEY) _enable_span.set_attribute(REDIRECT_ENABLE_KEY, True) _enable_span.end() @@ -89,7 +90,9 @@ async def send( raise exc break _redirect_span.set_attribute(REDIRECT_COUNT_KEY, len(history)) - new_request = self._build_redirect_request(request, response, current_options) + new_request = self._build_redirect_request( + request, response, current_options, request_options + ) history.append(request) request = new_request await response.aclose() @@ -111,7 +114,7 @@ def _get_current_options(self, request: httpx.Request) -> RedirectHandlerOption: Returns: RedirectHandlerOption: The options to used. """ - request_options = getattr(request, "options", None) + request_options = request.extensions.get(REQUEST_OPTIONS_KEY) if request_options: current_options = request_options.get( # type:ignore RedirectHandlerOption.get_key(), self.options) @@ -119,7 +122,11 @@ def _get_current_options(self, request: httpx.Request) -> RedirectHandlerOption: return self.options def _build_redirect_request( - self, request: httpx.Request, response: httpx.Response, options: RedirectHandlerOption + self, + request: httpx.Request, + response: httpx.Response, + options: RedirectHandlerOption, + request_options: typing.Optional[dict] = None, ) -> httpx.Request: """ Given a request and a redirect response, return a new request that @@ -129,13 +136,18 @@ def _build_redirect_request( url = self._redirect_url(request, response, options) stream = self._redirect_stream(request, method) + new_request_options = request_options.copy() if request_options else {} + extensions = request.extensions.copy() + if request_options is not None: + extensions[REQUEST_OPTIONS_KEY] = new_request_options + # Create the new request with the redirect URL and original headers new_request = httpx.Request( method=method, url=url, headers=request.headers.copy(), stream=stream, - extensions=request.extensions, + extensions=extensions, ) # Scrub sensitive headers before following the redirect @@ -154,7 +166,6 @@ def _build_redirect_request( if hasattr(request, "context"): new_request.context = request.context #type: ignore - new_request.options = {} #type: ignore return new_request def _redirect_method(self, request: httpx.Request, response: httpx.Response) -> str: diff --git a/packages/http/httpx/kiota_http/middleware/retry_handler.py b/packages/http/httpx/kiota_http/middleware/retry_handler.py index 37ef9b12..3078dba1 100644 --- a/packages/http/httpx/kiota_http/middleware/retry_handler.py +++ b/packages/http/httpx/kiota_http/middleware/retry_handler.py @@ -9,7 +9,7 @@ import httpx -from .middleware import BaseMiddleware +from .middleware import REQUEST_OPTIONS_KEY, BaseMiddleware from .options import RetryHandlerOption RETRY_ATTEMPT = "Retry-Attempt" @@ -105,9 +105,9 @@ def _get_current_options(self, request: httpx.Request) -> RetryHandlerOption: Returns: RetryHandlerOption: The options to used. """ - request_options = getattr(request, "options", None) + request_options = request.extensions.get(REQUEST_OPTIONS_KEY) if request_options: - current_options = request_options.get( # type:ignore + current_options = request_options.get( RetryHandlerOption.get_key(), self.options) return current_options return self.options diff --git a/packages/http/httpx/kiota_http/middleware/url_replace_handler.py b/packages/http/httpx/kiota_http/middleware/url_replace_handler.py index 3ad6aab0..a9f32f49 100644 --- a/packages/http/httpx/kiota_http/middleware/url_replace_handler.py +++ b/packages/http/httpx/kiota_http/middleware/url_replace_handler.py @@ -3,7 +3,7 @@ import httpx -from .middleware import BaseMiddleware +from .middleware import REQUEST_OPTIONS_KEY, BaseMiddleware from .options import UrlReplaceHandlerOption @@ -56,9 +56,9 @@ def _get_current_options(self, request: httpx.Request) -> UrlReplaceHandlerOptio Returns: UrlReplaceHandlerOption: The options to be used. """ - request_options = getattr(request, "options", None) + request_options = request.extensions.get(REQUEST_OPTIONS_KEY) if request_options: - current_options = request.options.get( # type:ignore + current_options = request_options.get( UrlReplaceHandlerOption.get_key(), self.options ) return current_options diff --git a/packages/http/httpx/kiota_http/middleware/user_agent_handler.py b/packages/http/httpx/kiota_http/middleware/user_agent_handler.py index 9ecc4779..c76dbc29 100644 --- a/packages/http/httpx/kiota_http/middleware/user_agent_handler.py +++ b/packages/http/httpx/kiota_http/middleware/user_agent_handler.py @@ -2,7 +2,7 @@ from httpx import AsyncBaseTransport, Request, Response -from .middleware import BaseMiddleware +from .middleware import REQUEST_OPTIONS_KEY, BaseMiddleware from .options import UserAgentHandlerOption @@ -39,9 +39,9 @@ def _get_current_options(self, request: Request) -> UserAgentHandlerOption: Returns: UserAgentHandlerOption: The options to be used. """ - request_options = getattr(request, "options", None) + request_options = request.extensions.get(REQUEST_OPTIONS_KEY) if request_options: - current_options = request.options.get( # type:ignore + current_options = request_options.get( UserAgentHandlerOption.get_key(), self.options ) return current_options diff --git a/packages/http/httpx/tests/middleware_tests/test_body_inspection_handler.py b/packages/http/httpx/tests/middleware_tests/test_body_inspection_handler.py new file mode 100644 index 00000000..e5c62965 --- /dev/null +++ b/packages/http/httpx/tests/middleware_tests/test_body_inspection_handler.py @@ -0,0 +1,657 @@ +import asyncio +import gzip +from io import BytesIO + +import pytest + +import httpx +import kiota_http.middleware.options.body_inspection_handler_option as option_module +from kiota_http.middleware import REQUEST_OPTIONS_KEY +from kiota_http.middleware.body_inspection_handler import BodyInspectionHandler +from kiota_http.middleware.options.body_inspection_handler_option import BodyInspectionHandlerOption +from kiota_http.middleware.redirect_handler import RedirectHandler + + +def test_default_options(): + """Ensures default values are disabled and bodies are None.""" + options = BodyInspectionHandlerOption() + assert not options.inspect_request_body + assert not options.inspect_response_body + assert options.request_body is None + assert options.response_body is None + assert options.get_request_body() is None + assert options.get_response_body() is None + assert options.get_request_body_stream() is None + assert options.get_response_body_stream() is None + assert options.request_body_stream is None + assert options.response_body_stream is None + + +def test_custom_options(): + """Ensures that custom boolean flags are properly set.""" + options = BodyInspectionHandlerOption( + inspect_request_body=True, + inspect_response_body=True, + ) + assert options.inspect_request_body + assert options.inspect_response_body + + +def test_options_stream_helpers_and_rewinding(): + """Ensures stream accessors return seekable streams rewound to position 0.""" + options = BodyInspectionHandlerOption( + request_body=b"request content", + response_body=b"response content", + ) + assert options.get_request_body() == b"request content" + assert options.get_response_body() == b"response content" + + req_stream1 = options.get_request_body_stream() + assert isinstance(req_stream1, BytesIO) + assert req_stream1.tell() == 0 + assert req_stream1.read() == b"request content" + + # A subsequent call returns a new rewound stream, not an exhausted one + req_stream2 = options.get_request_body_stream() + assert isinstance(req_stream2, BytesIO) + assert req_stream2.tell() == 0 + assert req_stream2.read() == b"request content" + + resp_stream1 = options.get_response_body_stream() + assert isinstance(resp_stream1, BytesIO) + assert resp_stream1.tell() == 0 + assert resp_stream1.read() == b"response content" + + resp_stream2 = options.response_body_stream + assert isinstance(resp_stream2, BytesIO) + assert resp_stream2.tell() == 0 + assert resp_stream2.read() == b"response content" + + +def test_handlers_do_not_share_options(): + """Two handlers without explicit options must not share option instances.""" + first = BodyInspectionHandler() + second = BodyInspectionHandler() + + assert first.options is not second.options + first.options.request_body = b"data" + assert second.options.request_body is None + + +def test_body_inspection_handler_construction(): + """Ensures BodyInspectionHandler can be constructed.""" + handler = BodyInspectionHandler() + assert handler is not None + assert isinstance(handler.options, BodyInspectionHandlerOption) + + +@pytest.mark.asyncio +async def test_rejects_null_request(): + """Ensures a null request produces an intentional error.""" + handler = BodyInspectionHandler() + transport = httpx.MockTransport(lambda _: httpx.Response(204)) + + with pytest.raises(TypeError, match="request cannot be null"): + await handler.send(None, transport) + + +@pytest.mark.asyncio +async def test_observability_span_covers_full_send(monkeypatch): + """Ensures telemetry covers the full body inspection handler execution.""" + attributes = {} + events = [] + + class RecordingSpan: + + def set_attribute(self, key, value): + attributes[key] = value + + def end(self): + events.append("span ended") + + async def response_body(): + events.append("response inspected") + yield b"response body" + + def request_handler(request: httpx.Request): + events.append("request sent") + return httpx.Response(200, content=response_body()) + + options = BodyInspectionHandlerOption(inspect_response_body=True) + handler = BodyInspectionHandler(options=options) + monkeypatch.setattr(handler, "_create_observability_span", lambda *_: RecordingSpan()) + + request = httpx.Request("GET", "https://localhost") + await handler.send(request, httpx.MockTransport(request_handler)) + + assert attributes == {"com.microsoft.kiota.handler.bodyInspection.enable": True} + assert events == ["request sent", "response inspected", "span ended"] + + +@pytest.mark.asyncio +async def test_observability_span_ends_when_send_raises(monkeypatch): + """Ensures telemetry ends when the downstream transport raises.""" + events = [] + + class RecordingSpan: + + def set_attribute(self, key, value): + pass + + def end(self): + events.append("span ended") + + def request_handler(request: httpx.Request): + events.append("request sent") + raise RuntimeError("transport failed") + + handler = BodyInspectionHandler() + monkeypatch.setattr(handler, "_create_observability_span", lambda *_: RecordingSpan()) + + request = httpx.Request("GET", "https://localhost") + transport = httpx.MockTransport(request_handler) + with pytest.raises(RuntimeError, match="transport failed"): + await handler.send(request, transport) + + assert events == ["request sent", "span ended"] + + +@pytest.mark.asyncio +async def test_inspect_request_body(): + """Ensures request body is captured and downstream transport still receives it.""" + received = [] + + def request_handler(request: httpx.Request): + received.append(request.read()) + return httpx.Response(200, json={"message": "ok"}) + + options = BodyInspectionHandlerOption(inspect_request_body=True) + handler = BodyInspectionHandler(options=options) + + request = httpx.Request("POST", "https://localhost", content=b"hello request") + mock_transport = httpx.MockTransport(request_handler) + + response = await handler.send(request, mock_transport) + + assert response.status_code == 200 + assert received == [b"hello request"] + assert handler.options.request_body == b"hello request" + assert handler.options.get_request_body() == b"hello request" + assert handler.options.get_request_body_stream().read() == b"hello request" + + +@pytest.mark.asyncio +async def test_inspect_response_body(): + """Ensures response body is captured and caller can still consume the response.""" + + def request_handler(request: httpx.Request): + return httpx.Response(200, content=b'{"user": "alice"}') + + options = BodyInspectionHandlerOption(inspect_response_body=True) + handler = BodyInspectionHandler(options=options) + + request = httpx.Request("GET", "https://localhost") + mock_transport = httpx.MockTransport(request_handler) + + response = await handler.send(request, mock_transport) + + assert response.status_code == 200 + # Downstream caller can still read the response content + assert response.content == b'{"user": "alice"}' + assert response.json() == {"user": "alice"} + # Handler option captured it + assert handler.options.response_body == b'{"user": "alice"}' + assert handler.options.get_response_body() == b'{"user": "alice"}' + assert handler.options.get_response_body_stream().read() == b'{"user": "alice"}' + + +@pytest.mark.asyncio +async def test_inspect_both_request_and_response_body(): + """Ensures both request and response bodies can be inspected simultaneously.""" + + def request_handler(request: httpx.Request): + return httpx.Response(201, content=b'{"created": true}') + + options = BodyInspectionHandlerOption( + inspect_request_body=True, + inspect_response_body=True, + ) + handler = BodyInspectionHandler(options=options) + + request = httpx.Request("POST", "https://localhost", content=b'{"name": "item"}') + mock_transport = httpx.MockTransport(request_handler) + + response = await handler.send(request, mock_transport) + + assert response.status_code == 201 + assert handler.options.request_body == b'{"name": "item"}' + assert handler.options.response_body == b'{"created": true}' + + +@pytest.mark.asyncio +async def test_disabled_inspection_does_not_capture(): + """Ensures bodies are not captured when inspection options are False.""" + + def request_handler(request: httpx.Request): + return httpx.Response(200, content=b'response content') + + handler = BodyInspectionHandler() + + request = httpx.Request("POST", "https://localhost", content=b'request content') + mock_transport = httpx.MockTransport(request_handler) + + response = await handler.send(request, mock_transport) + + assert response.status_code == 200 + assert handler.options.request_body is None + assert handler.options.response_body is None + + +@pytest.mark.asyncio +async def test_disabled_inspection_does_not_allocate_capture_state(): + """Ensures the default-disabled handler does not initialize context capture state.""" + token = option_module._BODY_CAPTURES.set(None) + try: + handler = BodyInspectionHandler() + + await handler.send( + httpx.Request("GET", "https://localhost"), + httpx.MockTransport(lambda _: httpx.Response(204)), + ) + + assert option_module._BODY_CAPTURES.get() is None + finally: + option_module._BODY_CAPTURES.reset(token) + + +@pytest.mark.asyncio +async def test_empty_bodies_returns_none(): + """Ensures empty request and response bodies result in None.""" + + def request_handler(request: httpx.Request): + return httpx.Response(204) + + options = BodyInspectionHandlerOption( + inspect_request_body=True, + inspect_response_body=True, + ) + handler = BodyInspectionHandler(options=options) + + request = httpx.Request("GET", "https://localhost") + mock_transport = httpx.MockTransport(request_handler) + + response = await handler.send(request, mock_transport) + + assert response.status_code == 204 + assert handler.options.request_body is None + assert handler.options.response_body is None + + +@pytest.mark.asyncio +async def test_streaming_payloads_inspection(): + """Ensures streaming generators for request and response are preserved and inspected.""" + received = [] + + async def req_gen(): + yield b"chunk1 " + yield b"chunk2" + + async def resp_gen(): + yield b"stream1 " + yield b"stream2" + + class RecordingTransport(httpx.AsyncBaseTransport): + + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + received.append(b"".join([chunk async for chunk in request.stream])) + return httpx.Response(200, content=resp_gen()) + + options = BodyInspectionHandlerOption( + inspect_request_body=True, + inspect_response_body=True, + ) + handler = BodyInspectionHandler(options=options) + + request = httpx.Request("POST", "https://localhost", content=req_gen()) + response = await handler.send(request, RecordingTransport()) + + assert response.status_code == 200 + assert received == [b"chunk1 chunk2"] + assert handler.options.request_body == b"chunk1 chunk2" + assert handler.options.response_body == b"stream1 stream2" + assert await response.aread() == b"stream1 stream2" + assert response.content == b"stream1 stream2" + + +@pytest.mark.asyncio +async def test_inspected_streaming_response_remains_raw_iterable(): + """Ensures inspection does not consume the response stream returned to the caller.""" + + async def resp_gen(): + yield b"stream1 " + yield b"stream2" + + def request_handler(request: httpx.Request): + return httpx.Response(200, content=resp_gen()) + + options = BodyInspectionHandlerOption(inspect_response_body=True) + handler = BodyInspectionHandler(options=options) + + response = await handler.send( + httpx.Request("GET", "https://localhost"), httpx.MockTransport(request_handler) + ) + assert response.num_bytes_downloaded == 0 + raw_content = b"".join([chunk async for chunk in response.aiter_raw()]) + + assert raw_content == b"stream1 stream2" + assert response.num_bytes_downloaded == len(raw_content) + assert options.response_body == b"stream1 stream2" + + +@pytest.mark.asyncio +async def test_inspected_streaming_response_remains_decoded_iterable(): + """Ensures decoded iteration consumes the restored stream and updates accounting.""" + decoded_content = b"decoded streaming response" + compressed_content = gzip.compress(decoded_content) + + async def resp_gen(): + yield compressed_content[:5] + yield compressed_content[5:] + + def request_handler(request: httpx.Request): + return httpx.Response( + 200, + headers={"Content-Encoding": "gzip"}, + content=resp_gen(), + ) + + options = BodyInspectionHandlerOption(inspect_response_body=True) + handler = BodyInspectionHandler(options=options) + + response = await handler.send( + httpx.Request("GET", "https://localhost"), httpx.MockTransport(request_handler) + ) + + assert options.response_body == decoded_content + assert not hasattr(response, "_content") + assert not response.is_stream_consumed + assert not response.is_closed + assert response.num_bytes_downloaded == 0 + + content = b"".join([chunk async for chunk in response.aiter_bytes()]) + + assert content == decoded_content + assert response.is_stream_consumed + assert response.is_closed + assert response.num_bytes_downloaded == len(compressed_content) + + +@pytest.mark.asyncio +async def test_already_consumed_uncached_response_does_not_fail_inspection(): + """Ensures an unrecoverable consumed stream is returned without inspection.""" + + async def resp_gen(): + yield b"already consumed" + + async def request_handler(request: httpx.Request): + response = httpx.Response(200, content=resp_gen(), request=request) + assert b"".join([chunk async for chunk in response.aiter_raw()]) == b"already consumed" + return response + + options = BodyInspectionHandlerOption(inspect_response_body=True) + handler = BodyInspectionHandler(options=options) + + response = await handler.send( + httpx.Request("GET", "https://localhost"), httpx.MockTransport(request_handler) + ) + + assert options.response_body is None + assert not hasattr(response, "_content") + assert response.is_stream_consumed + assert response.is_closed + + +@pytest.mark.asyncio +async def test_closed_unconsumed_response_does_not_fail_inspection(): + """Ensures an unreadable closed stream is returned without inspection.""" + + async def resp_gen(): + yield b"closed before consumption" + + async def request_handler(request: httpx.Request): + response = httpx.Response(200, content=resp_gen(), request=request) + await response.aclose() + return response + + options = BodyInspectionHandlerOption(inspect_response_body=True) + handler = BodyInspectionHandler(options=options) + + response = await handler.send( + httpx.Request("GET", "https://localhost"), httpx.MockTransport(request_handler) + ) + + assert options.response_body is None + assert not hasattr(response, "_content") + assert not response.is_stream_consumed + assert response.is_closed + + +@pytest.mark.asyncio +async def test_inspected_buffered_response_keeps_consumed_raw_stream_state(): + """Ensures decoded buffered content is not exposed as replayable raw bytes.""" + compressed_content = gzip.compress(b"decoded response") + + def request_handler(request: httpx.Request): + response = httpx.Response( + 200, + headers={"Content-Encoding": "gzip"}, + stream=httpx.ByteStream(compressed_content), + ) + response.read() + return response + + options = BodyInspectionHandlerOption(inspect_response_body=True) + handler = BodyInspectionHandler(options=options) + + response = await handler.send( + httpx.Request("GET", "https://localhost"), httpx.MockTransport(request_handler) + ) + + assert options.response_body == b"decoded response" + assert response.content == b"decoded response" + assert response.is_stream_consumed + assert response.is_closed + raw_stream = response.aiter_raw() + with pytest.raises(httpx.StreamConsumed): + await anext(raw_stream) + + +@pytest.mark.asyncio +async def test_per_request_options_override(): + """Ensures request-level options override handler-level options.""" + + def request_handler(request: httpx.Request): + return httpx.Response(200, content=b"response from server") + + # Handler defaults to inspection disabled + handler = BodyInspectionHandler() + + per_request_option = BodyInspectionHandlerOption( + inspect_request_body=True, + inspect_response_body=True, + ) + request = httpx.Request( + "POST", + "https://localhost", + content=b"request to server", + extensions={ + REQUEST_OPTIONS_KEY: { + BodyInspectionHandlerOption.get_key(): per_request_option, + } + }, + ) + + mock_transport = httpx.MockTransport(request_handler) + response = await handler.send(request, mock_transport) + + assert response.status_code == 200 + # Captured on per_request_option + assert per_request_option.request_body == b"request to server" + assert per_request_option.response_body == b"response from server" + # Handler options remained None + assert handler.options.request_body is None + assert handler.options.response_body is None + + +@pytest.mark.asyncio +async def test_per_request_options_apply_to_redirected_response(): + """Ensures redirect requests retain the original body inspection option.""" + + def request_handler(request: httpx.Request): + if request.url.path == "/redirected": + return httpx.Response(200, content=b"final response") + return httpx.Response( + 302, + headers={"Location": "/redirected"}, + content=b"redirect response", + ) + + redirect_handler = RedirectHandler() + redirect_handler.next = BodyInspectionHandler() + per_request_option = BodyInspectionHandlerOption(inspect_response_body=True) + request = httpx.Request( + "GET", + "https://localhost", + extensions={ + REQUEST_OPTIONS_KEY: { + BodyInspectionHandlerOption.get_key(): per_request_option, + } + }, + ) + + response = await redirect_handler.send(request, httpx.MockTransport(request_handler)) + + assert response.status_code == 200 + assert per_request_option.response_body == b"final response" + + +@pytest.mark.asyncio +async def test_reused_per_request_option_clears_previous_bodies(): + """Ensures disabled inspection does not retain captures from a previous request.""" + + def request_handler(request: httpx.Request): + return httpx.Response(200, content=b"response body") + + handler = BodyInspectionHandler() + option = BodyInspectionHandlerOption( + inspect_request_body=True, + inspect_response_body=True, + ) + first_request = httpx.Request( + "POST", + "https://localhost", + content=b"request body", + extensions={ + REQUEST_OPTIONS_KEY: { + BodyInspectionHandlerOption.get_key(): option, + } + }, + ) + transport = httpx.MockTransport(request_handler) + + await handler.send(first_request, transport) + assert option.request_body == b"request body" + assert option.response_body == b"response body" + + option.inspect_request_body = False + option.inspect_response_body = False + second_request = httpx.Request( + "GET", + "https://localhost", + extensions={ + REQUEST_OPTIONS_KEY: { + BodyInspectionHandlerOption.get_key(): option, + } + }, + ) + + await handler.send(second_request, transport) + + assert option.request_body is None + assert option.response_body is None + + +@pytest.mark.asyncio +async def test_handler_clears_body_on_subsequent_requests(): + """Ensures handler resets captured bodies across successive requests.""" + responses = [ + httpx.Response(200, content=b"response 1"), + httpx.Response(200, content=b"response 2"), + httpx.Response(204), + ] + + def request_handler(request: httpx.Request): + return responses.pop(0) + + options = BodyInspectionHandlerOption( + inspect_request_body=True, + inspect_response_body=True, + ) + handler = BodyInspectionHandler(options=options) + mock_transport = httpx.MockTransport(request_handler) + + # First request + req1 = httpx.Request("POST", "https://localhost", content=b"req 1") + await handler.send(req1, mock_transport) + assert handler.options.request_body == b"req 1" + assert handler.options.response_body == b"response 1" + + # Second request + req2 = httpx.Request("POST", "https://localhost", content=b"req 2") + await handler.send(req2, mock_transport) + assert handler.options.request_body == b"req 2" + assert handler.options.response_body == b"response 2" + + # Third request with no body + req3 = httpx.Request("GET", "https://localhost") + await handler.send(req3, mock_transport) + assert handler.options.request_body is None + assert handler.options.response_body is None + + +@pytest.mark.asyncio +async def test_concurrent_requests_keep_captures_isolated(): + """Ensures overlapping requests do not overwrite each other's captured bodies.""" + request_count = 0 + requests_ready = asyncio.Event() + + async def request_handler(request: httpx.Request): + nonlocal request_count + request_count += 1 + if request_count == 2: + requests_ready.set() + await requests_ready.wait() + content = await request.aread() + return httpx.Response(200, content=b"response: " + content) + + options = BodyInspectionHandlerOption( + inspect_request_body=True, + inspect_response_body=True, + ) + handler = BodyInspectionHandler(options=options) + transport = httpx.MockTransport(request_handler) + + async def send_and_capture(content: bytes): + request = httpx.Request("POST", "https://localhost", content=content) + await handler.send(request, transport) + return options.request_body, options.response_body + + captures = await asyncio.gather( + send_and_capture(b"request 1"), + send_and_capture(b"request 2"), + ) + + assert set(captures) == { + (b"request 1", b"response: request 1"), + (b"request 2", b"response: request 2"), + } diff --git a/packages/http/httpx/tests/test_httpx_request_adapter.py b/packages/http/httpx/tests/test_httpx_request_adapter.py index 1bef2d46..5b508d2e 100644 --- a/packages/http/httpx/tests/test_httpx_request_adapter.py +++ b/packages/http/httpx/tests/test_httpx_request_adapter.py @@ -14,6 +14,7 @@ from opentelemetry import trace from kiota_http.httpx_request_adapter import HttpxRequestAdapter +from kiota_http.middleware import REQUEST_OPTIONS_KEY from kiota_http.middleware.options import ResponseHandlerOption from .helpers import MockResponseObject @@ -124,6 +125,8 @@ def test_get_request_from_request_information(request_adapter, request_info, moc span = mock_otel_span req = request_adapter.get_request_from_request_information(request_info, span, span) assert isinstance(req, httpx.Request) + assert REQUEST_OPTIONS_KEY in req.extensions + assert req.extensions[REQUEST_OPTIONS_KEY] def test_get_response_handler(request_adapter, request_info): diff --git a/packages/http/httpx/tests/test_kiota_client_factory.py b/packages/http/httpx/tests/test_kiota_client_factory.py index b2cf5251..7dc98ca9 100644 --- a/packages/http/httpx/tests/test_kiota_client_factory.py +++ b/packages/http/httpx/tests/test_kiota_client_factory.py @@ -2,13 +2,22 @@ import httpx import pytest - from kiota_http.kiota_client_factory import KiotaClientFactory from kiota_http.middleware import ( - AsyncKiotaTransport, MiddlewarePipeline, ParametersNameDecodingHandler, RedirectHandler, - RetryHandler, UrlReplaceHandler, HeadersInspectionHandler + AsyncKiotaTransport, + BodyInspectionHandler, + HeadersInspectionHandler, + MiddlewarePipeline, + ParametersNameDecodingHandler, + RedirectHandler, + RetryHandler, + UrlReplaceHandler, +) +from kiota_http.middleware.options import ( + BodyInspectionHandlerOption, + RedirectHandlerOption, + RetryHandlerOption, ) -from kiota_http.middleware.options import RedirectHandlerOption, RetryHandlerOption from kiota_http.middleware.user_agent_handler import UserAgentHandler @@ -122,32 +131,37 @@ def test_get_default_middleware(): """Test fetching of default middleware with no custom options passed""" middleware = KiotaClientFactory.get_default_middleware(None) - assert len(middleware) == 6 + assert len(middleware) == 7 assert isinstance(middleware[0], RedirectHandler) assert isinstance(middleware[1], RetryHandler) assert isinstance(middleware[2], ParametersNameDecodingHandler) assert isinstance(middleware[3], UrlReplaceHandler) assert isinstance(middleware[4], UserAgentHandler) assert isinstance(middleware[5], HeadersInspectionHandler) + assert isinstance(middleware[6], BodyInspectionHandler) def test_get_default_middleware_with_options(): """Test fetching of default middleware with custom options passed""" retry_options = RetryHandlerOption(max_retries=7) redirect_options = RedirectHandlerOption(should_redirect=False) + body_inspection_options = BodyInspectionHandlerOption(inspect_request_body=True) options = { f'{retry_options.get_key()}': retry_options, - f'{redirect_options.get_key()}': redirect_options + f'{redirect_options.get_key()}': redirect_options, + f'{body_inspection_options.get_key()}': body_inspection_options, } middleware = KiotaClientFactory.get_default_middleware(options=options) - assert len(middleware) == 6 + assert len(middleware) == 7 assert isinstance(middleware[0], RedirectHandler) assert middleware[0].options.should_redirect is False assert isinstance(middleware[1], RetryHandler) assert middleware[1].options.max_retry == 7 assert isinstance(middleware[2], ParametersNameDecodingHandler) + assert isinstance(middleware[6], BodyInspectionHandler) + assert middleware[6].options.inspect_request_body is True def test_create_middleware_pipeline():