diff --git a/packages/http/httpx/kiota_http/httpx_request_adapter.py b/packages/http/httpx/kiota_http/httpx_request_adapter.py index a07c8757..20904ded 100644 --- a/packages/http/httpx/kiota_http/httpx_request_adapter.py +++ b/packages/http/httpx/kiota_http/httpx_request_adapter.py @@ -626,26 +626,29 @@ async def retry_cae_response_if_required( self, resp: httpx.Response, request_info: RequestInformation, claims: str ) -> httpx.Response: parent_span = self.start_tracing_span(request_info, "retry_cae_response_if_required") - if ( - resp.status_code == 401 - and not claims # previous claims exist. Means request has already been retried - and resp.headers.get(self.RESPONSE_AUTH_HEADER) - ): - auth_header_value = resp.headers.get(self.RESPONSE_AUTH_HEADER) - if auth_header_value.casefold().startswith( - self.BEARER_AUTHENTICATION_SCHEME.casefold() + try: + if ( + resp.status_code == 401 + and not claims # previous claims exist. Means request has already been retried + and resp.headers.get(self.RESPONSE_AUTH_HEADER) ): - claims_match = re.search('claims="([^"]+)"', auth_header_value) - if not claims_match: - return resp - response_claims = claims_match.group(1) - parent_span.add_event(AUTHENTICATE_CHALLENGED_EVENT_KEY) - parent_span.set_attribute("http.retry_count", 1) - return await self.get_http_response_message( - request_info, parent_span, response_claims - ) + auth_header_value = resp.headers.get(self.RESPONSE_AUTH_HEADER) + if auth_header_value.casefold().startswith( + self.BEARER_AUTHENTICATION_SCHEME.casefold() + ): + claims_match = re.search('claims="([^"]+)"', auth_header_value) + if not claims_match: + return resp + response_claims = claims_match.group(1) + parent_span.add_event(AUTHENTICATE_CHALLENGED_EVENT_KEY) + parent_span.set_attribute("http.retry_count", 1) + return await self.get_http_response_message( + request_info, parent_span, response_claims + ) + return resp return resp - return resp + finally: + parent_span.end() def get_response_handler(self, request_info: RequestInformation) -> Any: response_handler_option = request_info.request_options.get(ResponseHandlerOption.get_key()) diff --git a/packages/http/httpx/kiota_http/middleware/redirect_handler.py b/packages/http/httpx/kiota_http/middleware/redirect_handler.py index f0f8a930..9be19db8 100644 --- a/packages/http/httpx/kiota_http/middleware/redirect_handler.py +++ b/packages/http/httpx/kiota_http/middleware/redirect_handler.py @@ -74,27 +74,30 @@ async def send( _redirect_span = self._create_observability_span( request, f"RedirectHandler_send - redirect {len(history)}" ) - response = await super().send(request, transport) - _redirect_span.set_attribute(HTTP_RESPONSE_STATUS_CODE, response.status_code) - redirect_location = self.get_redirect_location(response) - - if redirect_location and current_options.should_redirect: - max_redirect -= 1 - if not self.increment(response, max_redirect, history[:]): - break - _redirect_span.set_attribute(REDIRECT_COUNT_KEY, len(history)) - new_request = self._build_redirect_request(request, response, current_options) - history.append(request) - request = new_request - await response.aclose() - continue - break + try: + response = await super().send(request, transport) + _redirect_span.set_attribute(HTTP_RESPONSE_STATUS_CODE, response.status_code) + redirect_location = self.get_redirect_location(response) + + if redirect_location and current_options.should_redirect: + max_redirect -= 1 + if not self.increment(response, max_redirect, history[:]): + if max_redirect < 0: + response.history = history + exc = RedirectError(f"Too many redirects. {response.history}") + _redirect_span.record_exception(exc) + raise exc + break + _redirect_span.set_attribute(REDIRECT_COUNT_KEY, len(history)) + new_request = self._build_redirect_request(request, response, current_options) + history.append(request) + request = new_request + await response.aclose() + continue + break + finally: + _redirect_span.end() response.history = history - if max_redirect < 0: - exc = RedirectError(f"Too many redirects. {response.history}") - _redirect_span.record_exception(exc) - _redirect_span.end() - raise exc return response diff --git a/packages/http/httpx/tests/conftest.py b/packages/http/httpx/tests/conftest.py index 55c9c124..db7668d1 100644 --- a/packages/http/httpx/tests/conftest.py +++ b/packages/http/httpx/tests/conftest.py @@ -8,12 +8,27 @@ from kiota_abstractions.method import Method from kiota_abstractions.request_information import RequestInformation from opentelemetry import trace +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter from kiota_http.httpx_request_adapter import HttpxRequestAdapter from .helpers import MockTransport, MockErrorObject, MockResponseObject, OfficeLocation +@pytest.fixture +def span_exporter(monkeypatch): + exporter = InMemorySpanExporter() + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + tracer = provider.get_tracer(__name__) + monkeypatch.setattr("kiota_http.middleware.middleware.tracer", tracer) + monkeypatch.setattr("kiota_http.httpx_request_adapter.tracer", tracer) + yield exporter + provider.shutdown() + + @pytest.fixture def sample_headers(): return {"Content-Type": "application/json"} diff --git a/packages/http/httpx/tests/middleware_tests/test_redirect_handler.py b/packages/http/httpx/tests/middleware_tests/test_redirect_handler.py index 4963af5e..232b7ea0 100644 --- a/packages/http/httpx/tests/middleware_tests/test_redirect_handler.py +++ b/packages/http/httpx/tests/middleware_tests/test_redirect_handler.py @@ -1,6 +1,9 @@ +import asyncio + import httpx import pytest +from kiota_http._exceptions import RedirectError from kiota_http.middleware import RedirectHandler from kiota_http.middleware.options import RedirectHandlerOption @@ -15,6 +18,70 @@ PERMANENT_REDIRECT = 308 +@pytest.mark.asyncio +@pytest.mark.parametrize("redirects, should_redirect", [(0, True), (2, True), (1, False)]) +async def test_redirect_spans_end(redirects, should_redirect, span_exporter): + requests = [] + + def request_handler(request): + requests.append(request) + if len(requests) <= redirects: + return httpx.Response(302, headers={LOCATION_HEADER: f"/redirect/{len(requests)}"}) + return httpx.Response(200) + + options = RedirectHandlerOption() + options.should_redirect = should_redirect + async with httpx.MockTransport(request_handler) as transport: + response = await RedirectHandler(options).send(httpx.Request("GET", BASE_URL), transport) + + attempts = redirects + 1 if should_redirect else 1 + assert len(requests) == attempts + assert response.status_code == (200 if should_redirect else 302) + assert len(response.history) == attempts - 1 + spans = span_exporter.get_finished_spans() + assert [span.name for span in spans] == ["RedirectHandler_send"] + [ + f"RedirectHandler_send - redirect {index}" for index in range(attempts) + ] + assert spans[-1].attributes["http.response.status_code"] == response.status_code + + +@pytest.mark.asyncio +async def test_redirect_limit_span_records_error_before_ending(span_exporter): + options = RedirectHandlerOption() + options.max_redirect = 1 + transport = httpx.MockTransport( + lambda request: httpx.Response(302, headers={LOCATION_HEADER: "/next"}) + ) + async with transport: + with pytest.raises(RedirectError, match="Too many redirects") as error: + await RedirectHandler(options).send(httpx.Request("GET", BASE_URL), transport) + + spans = span_exporter.get_finished_spans() + assert [span.name for span in spans] == [ + "RedirectHandler_send", "RedirectHandler_send - redirect 0", + "RedirectHandler_send - redirect 1" + ] + assert spans[-1].events[0].name == "exception" + assert spans[-1].events[0].attributes["exception.message"] == str(error.value) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("error_type", [httpx.ReadError, asyncio.CancelledError]) +async def test_redirect_span_ends_on_transport_failure(error_type, span_exporter): + error = error_type("request interrupted") + + def request_handler(request): + raise error + + async with httpx.MockTransport(request_handler) as transport: + with pytest.raises(error_type) as raised: + await RedirectHandler().send(httpx.Request("GET", BASE_URL), transport) + assert raised.value is error + assert [span.name for span in span_exporter.get_finished_spans()] == [ + "RedirectHandler_send", "RedirectHandler_send - redirect 0" + ] + + @pytest.fixture def mock_redirect_handler(): return RedirectHandler() diff --git a/packages/http/httpx/tests/test_httpx_request_adapter.py b/packages/http/httpx/tests/test_httpx_request_adapter.py index a2e50852..1bef2d46 100644 --- a/packages/http/httpx/tests/test_httpx_request_adapter.py +++ b/packages/http/httpx/tests/test_httpx_request_adapter.py @@ -22,6 +22,59 @@ BASE_URL = "https://graph.microsoft.com" +@pytest.mark.asyncio +@pytest.mark.parametrize("status, header, claims", [ + (200, None, ""), + (401, None, ""), + (401, 'Bearer claims="challenge"', "previous-claims"), + (401, "Basic realm=test", ""), + (401, "Bearer", ""), +]) +async def test_cae_span_ends_without_retry( + status, header, claims, request_adapter, request_info, span_exporter +): + headers = {"WWW-Authenticate": header} if header else {} + response = httpx.Response(status, headers=headers) + request_adapter.get_http_response_message = AsyncMock() + result = await request_adapter.retry_cae_response_if_required(response, request_info, claims) + assert result is response + request_adapter.get_http_response_message.assert_not_awaited() + spans = span_exporter.get_finished_spans() + assert [span.name for span in spans] == ["retry_cae_response_if_required - UNKNOWN"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("error_type", [None, httpx.ReadError, asyncio.CancelledError]) +async def test_cae_span_ends_after_retry( + error_type, request_adapter, request_info, span_exporter +): + response = httpx.Response(401, headers={"WWW-Authenticate": 'Bearer claims="challenge"'}) + retried_response = httpx.Response(200) + error = error_type("retry interrupted") if error_type else None + + async def retry(info, span, claims): + assert info is request_info + assert claims == "challenge" + assert span.is_recording() + if error is not None: + raise error + return retried_response + + request_adapter.get_http_response_message = AsyncMock(side_effect=retry) + if error_type: + with pytest.raises(error_type) as raised: + await request_adapter.retry_cae_response_if_required(response, request_info, "") + assert raised.value is error + else: + result = await request_adapter.retry_cae_response_if_required(response, request_info, "") + assert result is retried_response + request_adapter.get_http_response_message.assert_awaited_once() + spans = span_exporter.get_finished_spans() + assert [span.name for span in spans] == ["retry_cae_response_if_required - UNKNOWN"] + assert spans[0].attributes["http.retry_count"] == 1 + assert spans[0].events[0].name == "com.microsoft.kiota.authenticate_challenge_received" + + def test_create_request_adapter(auth_provider): request_adapter = HttpxRequestAdapter(auth_provider) assert request_adapter._authentication_provider is auth_provider @@ -420,7 +473,7 @@ async def test_observability( @pytest.mark.asyncio async def test_retries_on_cae_failure( - request_adapter, request_info_mock, mock_cae_failure_response, mock_otel_span + request_adapter, request_info_mock, mock_cae_failure_response, mock_otel_span, span_exporter ): request_adapter._http_client.send = AsyncMock(return_value=mock_cae_failure_response) request_adapter._authentication_provider.authenticate_request = AsyncMock() @@ -439,6 +492,13 @@ async def test_retries_on_cae_failure( ), ] request_adapter._authentication_provider.authenticate_request.assert_has_awaits(calls) + cae_spans = [ + span for span in span_exporter.get_finished_spans() + if span.name.startswith("retry_cae_response_if_required - ") + ] + assert len(cae_spans) == 2 + assert cae_spans[0].attributes.get("http.retry_count") is None + assert cae_spans[1].attributes["http.retry_count"] == 1 @pytest.mark.asyncio