diff --git a/src/quilt_hp/transport.py b/src/quilt_hp/transport.py index 2cff92a..7172049 100644 --- a/src/quilt_hp/transport.py +++ b/src/quilt_hp/transport.py @@ -135,7 +135,7 @@ async def intercept_unary_unary( return await cast("Awaitable[object]", call) except grpc.aio.AioRpcError as exc: if exc.code() == grpc.StatusCode.UNAUTHENTICATED and self._refresh_callback: - logger.warning("Retrying unary RPC after UNAUTHENTICATED response") + logger.info("Retrying unary RPC after UNAUTHENTICATED response") return await self._refresh_and_retry_unary( continuation, client_call_details, request ) @@ -155,7 +155,7 @@ async def intercept_unary_stream( await cast("Any", call).wait_for_connection() except grpc.aio.AioRpcError as exc: if exc.code() == grpc.StatusCode.UNAUTHENTICATED and self._refresh_callback: - logger.warning("Retrying streaming RPC setup after UNAUTHENTICATED response") + logger.info("Retrying streaming RPC setup after UNAUTHENTICATED response") await self._refresh() retried = await continuation(self._patch(client_call_details), request) try: diff --git a/tests/test_transport_interceptor_extra.py b/tests/test_transport_interceptor_extra.py index 0bebe76..1767e2e 100644 --- a/tests/test_transport_interceptor_extra.py +++ b/tests/test_transport_interceptor_extra.py @@ -1,6 +1,7 @@ from __future__ import annotations import inspect +import logging from unittest.mock import MagicMock import grpc @@ -136,6 +137,87 @@ async def _continuation(call_details: grpc.aio.ClientCallDetails, request: objec assert refreshed == ["yes"] +@pytest.mark.asyncio +async def test_auth_interceptor_logs_unary_retry_at_info( + caplog: pytest.LogCaptureFixture, +) -> None: + refreshed: list[str] = [] + + async def _refresh(_context: transport.TokenRefreshContext) -> None: + refreshed.append("yes") + + interceptor = transport._AuthInterceptor(lambda: "******", refresh_callback=_refresh) + details = grpc.aio.ClientCallDetails( + method="/svc/method", + timeout=1, + metadata=None, + credentials=None, + wait_for_ready=False, + ) + + calls = 0 + + async def _continuation(_call_details: grpc.aio.ClientCallDetails, request: object) -> object: + nonlocal calls + calls += 1 + if calls == 1: + return _FakeCall(error=_FakeRpcError(grpc.StatusCode.UNAUTHENTICATED, "expired")) + return _FakeCall(result=request) + + with caplog.at_level(logging.INFO): + assert await interceptor.intercept_unary_unary(_continuation, details, "req") == "req" + + matching = [ + record + for record in caplog.records + if record.getMessage() == "Retrying unary RPC after UNAUTHENTICATED response" + ] + assert matching + assert all(record.levelno == logging.INFO for record in matching) + assert refreshed == ["yes"] + + +@pytest.mark.asyncio +async def test_auth_interceptor_logs_stream_setup_retry_at_info( + caplog: pytest.LogCaptureFixture, +) -> None: + refreshed: list[str] = [] + + async def _refresh(_context: transport.TokenRefreshContext) -> None: + refreshed.append("yes") + + interceptor = transport._AuthInterceptor(lambda: "******", refresh_callback=_refresh) + details = grpc.aio.ClientCallDetails( + method="/svc/method", + timeout=1, + metadata=None, + credentials=None, + wait_for_ready=False, + ) + + calls = 0 + + async def _continuation(_call_details: grpc.aio.ClientCallDetails, request: object) -> object: + nonlocal calls + calls += 1 + if calls == 1: + return _FakeCall(error=_FakeRpcError(grpc.StatusCode.UNAUTHENTICATED, "expired")) + return _FakeCall(result=request) + + with caplog.at_level(logging.INFO): + stream_call = await interceptor.intercept_unary_stream(_continuation, details, "req") + + matching = [ + record + for record in caplog.records + if record.getMessage() == "Retrying streaming RPC setup after UNAUTHENTICATED response" + ] + assert matching + assert all(record.levelno == logging.INFO for record in matching) + assert await stream_call == "req" # type: ignore[misc] + assert refreshed == ["yes"] + + @pytest.mark.asyncio async def test_auth_interceptor_non_retry_paths() -> None: interceptor = transport._AuthInterceptor(lambda: "Bearer abc")