diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 6f2d7ee5921..871f4475a0b 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -15,7 +15,9 @@ from fastapi import Request, UploadFile from starlette.datastructures import FormData, Headers, QueryParams from starlette.datastructures import UploadFile as StarletteUploadFile -sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system path +sys.path.insert( + 0, os.path.abspath("../../..") +) # Adds the parent directory to the system path from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS, @@ -65,7 +67,9 @@ async def test_build_request_files_from_upload_file(): upload_file = UploadFile(file=file, filename="test.txt", headers=headers) upload_file.read = AsyncMock(return_value=file_content) - result = await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file(upload_file) + result = await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file( + upload_file + ) assert result == ("test.txt", file_content, "text/plain") # Test with Starlette UploadFile @@ -77,7 +81,9 @@ async def test_build_request_files_from_upload_file(): ) starlette_file.read = AsyncMock(return_value=file_content) - result = await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file(starlette_file) + result = await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file( + starlette_file + ) assert result == ("test2.txt", file_content, "text/plain") @@ -263,7 +269,9 @@ async def test_non_streaming_http_request_handler_multipart_with_non_empty_parse """ request = MagicMock(spec=Request) request.method = "POST" - request.headers = Headers({"content-type": "multipart/form-data; boundary=------------------------test"}) + request.headers = Headers( + {"content-type": "multipart/form-data; boundary=------------------------test"} + ) file_content = b"test file content" file = BytesIO(file_content) @@ -302,7 +310,9 @@ async def test_pass_through_request_failure_handler(): Critical Test: When a users pass through endpoint request fails, we must log the failure code, exception in litellm spend logs. """ with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: - with patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client") as mock_get_client: + with patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client" + ) as mock_get_client: with patch( "litellm.proxy.pass_through_endpoints.pass_through_endpoints.ProxyBaseLLMRequestProcessing" ) as mock_processing: @@ -313,7 +323,9 @@ async def test_pass_through_request_failure_handler(): # Setup mock for httpx client mock_client = MagicMock() mock_client.client = MagicMock() - mock_client.client.request = AsyncMock(side_effect=httpx.HTTPError("Request failed")) + mock_client.client.request = AsyncMock( + side_effect=httpx.HTTPError("Request failed") + ) mock_get_client.return_value = mock_client # Mock headers for custom headers @@ -346,7 +358,9 @@ async def test_pass_through_request_failure_handler(): # Verify the arguments to post_call_failure_hook call_args = mock_proxy_logging.post_call_failure_hook.call_args[1] assert call_args["user_api_key_dict"] == mock_user_api_key_dict - assert isinstance(call_args["original_exception"], TypeError) # Now expecting TypeError + assert isinstance( + call_args["original_exception"], TypeError + ) # Now expecting TypeError assert "traceback_str" in call_args @@ -357,14 +371,27 @@ def test_is_langfuse_route(): handler = PassThroughEndpointLogging() # Test positive cases - assert handler.is_langfuse_route("http://localhost:4000/langfuse/api/public/traces") is True - assert handler.is_langfuse_route("https://proxy.example.com/langfuse/api/public/sessions") is True + assert ( + handler.is_langfuse_route("http://localhost:4000/langfuse/api/public/traces") + is True + ) + assert ( + handler.is_langfuse_route( + "https://proxy.example.com/langfuse/api/public/sessions" + ) + is True + ) assert handler.is_langfuse_route("/langfuse/api/public/ingestion") is True assert handler.is_langfuse_route("http://localhost:4000/langfuse/") is True # Test negative cases - assert handler.is_langfuse_route("https://api.openai.com/v1/chat/completions") is False - assert handler.is_langfuse_route("http://localhost:4000/anthropic/v1/messages") is False + assert ( + handler.is_langfuse_route("https://api.openai.com/v1/chat/completions") is False + ) + assert ( + handler.is_langfuse_route("http://localhost:4000/anthropic/v1/messages") + is False + ) assert handler.is_langfuse_route("https://example.com/other") is False assert handler.is_langfuse_route("") is False @@ -418,7 +445,10 @@ async def test_langfuse_passthrough_no_logging(): assert result is None # Verify that the passthrough_logging_payload was still set (this happens before the langfuse check) - assert mock_logging_obj.model_call_details["passthrough_logging_payload"] == passthrough_logging_payload + assert ( + mock_logging_obj.model_call_details["passthrough_logging_payload"] + == passthrough_logging_payload + ) def test_construct_target_url_with_subpath(): @@ -866,7 +896,9 @@ async def test_create_pass_through_route_with_cost_per_request(): # Mock the pass_through_request function to capture its call with ( - patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_request") as mock_pass_through, + patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_request" + ) as mock_pass_through, patch( "litellm.proxy.pass_through_endpoints.pass_through_endpoints.InitPassThroughEndpointHelpers.is_registered_pass_through_route" ) as mock_is_registered, @@ -913,7 +945,10 @@ def test_resolve_pass_through_request_timeout_precedence(): assert resolve_pass_through_request_timeout(endpoint_timeout=800) == 800.0 with patch("litellm.proxy.proxy_server.general_settings", {}): - assert resolve_pass_through_request_timeout() == DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS + assert ( + resolve_pass_through_request_timeout() + == DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS + ) def test_resolve_llm_passthrough_timeout_precedence(): @@ -948,11 +983,15 @@ async def test_pass_through_request_uses_resolved_timeout(): with patch( "litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client" ) as mock_get_client: - mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda **kwargs: kwargs["data"]) + mock_proxy_logging.pre_call_hook = AsyncMock( + side_effect=lambda **kwargs: kwargs["data"] + ) mock_client = MagicMock() mock_client.client = MagicMock() - mock_client.client.request = AsyncMock(side_effect=httpx.HTTPError("Request failed")) + mock_client.client.request = AsyncMock( + side_effect=httpx.HTTPError("Request failed") + ) mock_get_client.return_value = mock_client mock_request = MagicMock(spec=Request) @@ -990,7 +1029,9 @@ async def test_create_pass_through_route_forwards_timeout(): ) with ( - patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_request") as mock_pass_through, + patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_request" + ) as mock_pass_through, patch( "litellm.proxy.pass_through_endpoints.pass_through_endpoints.InitPassThroughEndpointHelpers.is_registered_pass_through_route" ) as mock_is_registered, @@ -1103,7 +1144,9 @@ async def test_pass_through_request_contains_proxy_server_request_in_kwargs(): "litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_response_body" ) as mock_get_response_body: # Setup mock for pre_call_hook and post_call_failure_hook - mock_proxy_logging.pre_call_hook = AsyncMock(return_value={"test": "data"}) + mock_proxy_logging.pre_call_hook = AsyncMock( + return_value={"test": "data"} + ) mock_proxy_logging.post_call_failure_hook = AsyncMock() mock_proxy_logging.post_call_response_headers_hook = AsyncMock( return_value={"x-callback-test": "value"} @@ -1113,7 +1156,9 @@ async def test_pass_through_request_contains_proxy_server_request_in_kwargs(): mock_response = MagicMock() mock_response.status_code = 200 mock_response.headers = {} - mock_response.aread = AsyncMock(return_value=b'{"success": true}') + mock_response.aread = AsyncMock( + return_value=b'{"success": true}' + ) mock_response.text = '{"success": true}' mock_response.raise_for_status = MagicMock() @@ -1133,7 +1178,9 @@ async def test_pass_through_request_contains_proxy_server_request_in_kwargs(): mock_request = MagicMock(spec=Request) mock_request.method = "POST" mock_request.url = "http://test-proxy.com/api/endpoint" - mock_request.body = AsyncMock(return_value=b'{"message": "test request"}') + mock_request.body = AsyncMock( + return_value=b'{"message": "test request"}' + ) mock_request.headers = Headers({}) mock_request.query_params = QueryParams({}) @@ -1212,7 +1259,9 @@ async def test_pass_through_request_streaming_marks_logging_obj_as_stream(): with patch( "litellm.proxy.pass_through_endpoints.pass_through_endpoints.PassThroughStreamingHandler.chunk_processor" ) as mock_chunk_processor: - mock_proxy_logging.pre_call_hook = AsyncMock(return_value={"model": "claude-3", "stream": True}) + mock_proxy_logging.pre_call_hook = AsyncMock( + return_value={"model": "claude-3", "stream": True} + ) mock_proxy_logging.post_call_failure_hook = AsyncMock() mock_proxy_logging.post_call_response_headers_hook = AsyncMock( return_value={"x-callback-test": "value"} @@ -1237,7 +1286,9 @@ async def test_pass_through_request_streaming_marks_logging_obj_as_stream(): mock_request = MagicMock(spec=Request) mock_request.method = "POST" mock_request.url = "http://test-proxy.com/v1/messages" - mock_request.body = AsyncMock(return_value=b'{"model": "claude-3", "stream": true}') + mock_request.body = AsyncMock( + return_value=b'{"model": "claude-3", "stream": true}' + ) mock_request.headers = Headers({}) mock_request.query_params = QueryParams({}) @@ -1253,7 +1304,9 @@ async def test_pass_through_request_streaming_marks_logging_obj_as_stream(): assert async_client.send.call_args.kwargs["stream"] is True mock_chunk_processor.assert_called_once() - logging_obj = mock_chunk_processor.call_args.kwargs["litellm_logging_obj"] + logging_obj = mock_chunk_processor.call_args.kwargs[ + "litellm_logging_obj" + ] assert logging_obj.stream is True assert logging_obj.model_call_details["stream"] is True @@ -1274,7 +1327,9 @@ async def test_pass_through_request_sse_response_marks_logging_obj_as_stream(): with patch( "litellm.proxy.pass_through_endpoints.pass_through_endpoints.PassThroughStreamingHandler.chunk_processor" ) as mock_chunk_processor: - mock_proxy_logging.pre_call_hook = AsyncMock(return_value={"model": "claude-3"}) + mock_proxy_logging.pre_call_hook = AsyncMock( + return_value={"model": "claude-3"} + ) mock_proxy_logging.post_call_failure_hook = AsyncMock() mock_proxy_logging.post_call_response_headers_hook = AsyncMock( return_value={"x-callback-test": "value"} @@ -1314,7 +1369,9 @@ async def test_pass_through_request_sse_response_marks_logging_obj_as_stream(): async_client.send.assert_awaited_once() mock_chunk_processor.assert_called_once() - logging_obj = mock_chunk_processor.call_args.kwargs["litellm_logging_obj"] + logging_obj = mock_chunk_processor.call_args.kwargs[ + "litellm_logging_obj" + ] assert logging_obj.stream is True assert logging_obj.model_call_details["stream"] is True @@ -1341,10 +1398,16 @@ async def test_create_pass_through_endpoint(): ) # Mock the database functions - with patch("litellm.proxy.proxy_server.get_config_general_settings") as mock_get_config: - with patch("litellm.proxy.proxy_server.update_config_general_settings") as mock_update_config: + with patch( + "litellm.proxy.proxy_server.get_config_general_settings" + ) as mock_get_config: + with patch( + "litellm.proxy.proxy_server.update_config_general_settings" + ) as mock_update_config: # Mock existing config (empty list) - mock_get_config.return_value = ConfigFieldInfo(field_name="pass_through_endpoints", field_value=[]) + mock_get_config.return_value = ConfigFieldInfo( + field_name="pass_through_endpoints", field_value=[] + ) # Create test endpoint data test_endpoint = PassThroughGenericEndpoint( @@ -1414,8 +1477,12 @@ async def test_update_pass_through_endpoint(): ) # Mock the database functions - with patch("litellm.proxy.proxy_server.get_config_general_settings") as mock_get_config: - with patch("litellm.proxy.proxy_server.update_config_general_settings") as mock_update_config: + with patch( + "litellm.proxy.proxy_server.get_config_general_settings" + ) as mock_get_config: + with patch( + "litellm.proxy.proxy_server.update_config_general_settings" + ) as mock_update_config: # Create existing endpoint data existing_endpoint_id = "test-endpoint-123" existing_endpoints = [ @@ -1512,14 +1579,18 @@ async def test_create_pass_through_endpoint_auth_true_enforces_allowlist(): registry: dict = {} with ( - patch("litellm.proxy.proxy_server.get_config_general_settings") as mock_get_config, + patch( + "litellm.proxy.proxy_server.get_config_general_settings" + ) as mock_get_config, patch("litellm.proxy.proxy_server.update_config_general_settings"), patch( "litellm.proxy.pass_through_endpoints.pass_through_endpoints._registered_pass_through_routes", registry, ), ): - mock_get_config.return_value = ConfigFieldInfo(field_name="pass_through_endpoints", field_value=[]) + mock_get_config.return_value = ConfigFieldInfo( + field_name="pass_through_endpoints", field_value=[] + ) # auth is not passed -> defaults to True on PassThroughGenericEndpoint endpoint = PassThroughGenericEndpoint( @@ -1534,12 +1605,19 @@ async def test_create_pass_through_endpoint_auth_true_enforces_allowlist(): ) assert any(value.get("auth") is True for value in registry.values()) - assert RouteChecks.is_auth_enforced_pass_through_route(route="/secure-passthrough", method="POST") is True + assert ( + RouteChecks.is_auth_enforced_pass_through_route( + route="/secure-passthrough", method="POST" + ) + is True + ) post_request = MagicMock(spec=Request) post_request.method = "POST" - without_allowlist = UserAPIKeyAuth(user_id="u", allowed_routes=["llm_api_routes"]) + without_allowlist = UserAPIKeyAuth( + user_id="u", allowed_routes=["llm_api_routes"] + ) with pytest.raises(HTTPException) as exc_info: RouteChecks.is_virtual_key_allowed_to_call_route( route="/secure-passthrough", @@ -1596,7 +1674,9 @@ async def test_update_pass_through_endpoint_auth_true_enforces_allowlist(): ] with ( - patch("litellm.proxy.proxy_server.get_config_general_settings") as mock_get_config, + patch( + "litellm.proxy.proxy_server.get_config_general_settings" + ) as mock_get_config, patch("litellm.proxy.proxy_server.update_config_general_settings"), patch( "litellm.proxy.pass_through_endpoints.pass_through_endpoints._registered_pass_through_routes", @@ -1619,12 +1699,19 @@ async def test_update_pass_through_endpoint_auth_true_enforces_allowlist(): user_api_key_dict=MagicMock(spec=UserAPIKeyAuth), ) - assert RouteChecks.is_auth_enforced_pass_through_route(route="/edited-passthrough", method="POST") is True + assert ( + RouteChecks.is_auth_enforced_pass_through_route( + route="/edited-passthrough", method="POST" + ) + is True + ) post_request = MagicMock(spec=Request) post_request.method = "POST" - without_allowlist = UserAPIKeyAuth(user_id="u", allowed_routes=["llm_api_routes"]) + without_allowlist = UserAPIKeyAuth( + user_id="u", allowed_routes=["llm_api_routes"] + ) with pytest.raises(HTTPException) as exc_info: RouteChecks.is_virtual_key_allowed_to_call_route( route="/edited-passthrough", @@ -1666,8 +1753,12 @@ async def test_update_pass_through_endpoint_preserves_auth_false(): ] with ( - patch("litellm.proxy.proxy_server.get_config_general_settings") as mock_get_config, - patch("litellm.proxy.proxy_server.update_config_general_settings") as mock_update_config, + patch( + "litellm.proxy.proxy_server.get_config_general_settings" + ) as mock_get_config, + patch( + "litellm.proxy.proxy_server.update_config_general_settings" + ) as mock_update_config, patch( "litellm.proxy.pass_through_endpoints.pass_through_endpoints._registered_pass_through_routes", registry, @@ -1694,7 +1785,12 @@ async def test_update_pass_through_endpoint_preserves_auth_false(): persisted = mock_update_config.call_args[1]["data"].field_value[0] assert persisted["auth"] is False - assert RouteChecks.is_auth_enforced_pass_through_route(route="/public-passthrough", method="POST") is False + assert ( + RouteChecks.is_auth_enforced_pass_through_route( + route="/public-passthrough", method="POST" + ) + is False + ) @pytest.mark.asyncio @@ -1714,7 +1810,9 @@ async def test_update_pass_through_endpoint_not_found(): ) # Mock the database functions - with patch("litellm.proxy.proxy_server.get_config_general_settings") as mock_get_config: + with patch( + "litellm.proxy.proxy_server.get_config_general_settings" + ) as mock_get_config: # Mock existing config with different endpoint existing_endpoints = [ { @@ -1732,7 +1830,9 @@ async def test_update_pass_through_endpoint_not_found(): ) # Create update data - update_data = PassThroughGenericEndpoint(path="/test/endpoint", target="http://newapi.com/v2") + update_data = PassThroughGenericEndpoint( + path="/test/endpoint", target="http://newapi.com/v2" + ) # Mock user API key dict mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) @@ -1771,8 +1871,12 @@ async def test_delete_pass_through_endpoint(): ) # Mock the database functions - with patch("litellm.proxy.proxy_server.get_config_general_settings") as mock_get_config: - with patch("litellm.proxy.proxy_server.update_config_general_settings") as mock_update_config: + with patch( + "litellm.proxy.proxy_server.get_config_general_settings" + ) as mock_get_config: + with patch( + "litellm.proxy.proxy_server.update_config_general_settings" + ) as mock_update_config: # Create existing endpoint data endpoint_to_delete_id = "test-endpoint-123" other_endpoint_id = "other-endpoint-456" @@ -1850,7 +1954,9 @@ async def test_delete_pass_through_endpoint_not_found(): ) # Mock the database functions - with patch("litellm.proxy.proxy_server.get_config_general_settings") as mock_get_config: + with patch( + "litellm.proxy.proxy_server.get_config_general_settings" + ) as mock_get_config: # Mock existing config with different endpoint existing_endpoints = [ { @@ -1941,8 +2047,14 @@ async def test_get_pass_through_endpoints_includes_config_and_db(): with patch( "litellm.proxy.pass_through_endpoints.pass_through_endpoints._get_pass_through_endpoints_from_config" ) as mock_get_config: - db_objects = [PassThroughGenericEndpoint(**ep, is_from_config=False) for ep in db_endpoints] - config_objects = [PassThroughGenericEndpoint(**ep, is_from_config=True) for ep in config_endpoints] + db_objects = [ + PassThroughGenericEndpoint(**ep, is_from_config=False) + for ep in db_endpoints + ] + config_objects = [ + PassThroughGenericEndpoint(**ep, is_from_config=True) + for ep in config_endpoints + ] mock_get_db.return_value = db_objects mock_get_config.return_value = config_objects @@ -2016,9 +2128,13 @@ async def test_delete_pass_through_endpoint_empty_list(): ) # Mock the database functions - with patch("litellm.proxy.proxy_server.get_config_general_settings") as mock_get_config: + with patch( + "litellm.proxy.proxy_server.get_config_general_settings" + ) as mock_get_config: # Mock empty config - mock_get_config.return_value = ConfigFieldInfo(field_name="pass_through_endpoints", field_value=None) + mock_get_config.return_value = ConfigFieldInfo( + field_name="pass_through_endpoints", field_value=None + ) # Mock user API key dict mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) @@ -2057,7 +2173,9 @@ async def test_pass_through_request_query_params_forwarding(): ) as mock_get_response_body: # Setup mock for pre_call_hook test_body = {"name": "Azure Assistant", "model": "gpt-4o"} - mock_proxy_logging.pre_call_hook = AsyncMock(return_value=test_body) + mock_proxy_logging.pre_call_hook = AsyncMock( + return_value=test_body + ) mock_proxy_logging.post_call_response_headers_hook = AsyncMock( return_value={"x-callback-test": "value"} ) @@ -2066,7 +2184,9 @@ async def test_pass_through_request_query_params_forwarding(): mock_response = MagicMock() mock_response.status_code = 200 mock_response.headers = {"content-type": "application/json"} - mock_response.aread = AsyncMock(return_value=b'{"id": "asst_123", "object": "assistant"}') + mock_response.aread = AsyncMock( + return_value=b'{"id": "asst_123", "object": "assistant"}' + ) mock_response.text = '{"id": "asst_123", "object": "assistant"}' mock_response.raise_for_status = MagicMock() @@ -2088,12 +2208,20 @@ async def test_pass_through_request_query_params_forwarding(): # Create mock request with query parameters (Azure API version) mock_request = MagicMock(spec=Request) mock_request.method = "POST" - mock_request.url = "http://localhost:4000/azure-assistant/openai/assistants" - mock_request.body = AsyncMock(return_value=json.dumps(test_body).encode()) - mock_request.headers = Headers({"Content-Type": "application/json"}) + mock_request.url = ( + "http://localhost:4000/azure-assistant/openai/assistants" + ) + mock_request.body = AsyncMock( + return_value=json.dumps(test_body).encode() + ) + mock_request.headers = Headers( + {"Content-Type": "application/json"} + ) # Create QueryParams with api-version parameter - mock_request.query_params = QueryParams([("api-version", "2025-01-01-preview")]) + mock_request.query_params = QueryParams( + [("api-version", "2025-01-01-preview")] + ) # Create mock user API key dict mock_user_api_key_dict = MagicMock() @@ -2115,7 +2243,9 @@ async def test_pass_through_request_query_params_forwarding(): # The key assertion: query parameters should be preserved and passed to the HTTP handler assert "requested_query_params" in call_kwargs - assert call_kwargs["requested_query_params"] == {"api-version": "2025-01-01-preview"} + assert call_kwargs["requested_query_params"] == { + "api-version": "2025-01-01-preview" + } assert call_kwargs.get("forward_multipart") is False # Verify the target URL is correct @@ -2165,7 +2295,9 @@ async def _run_pass_through_and_capture_wire_url( "PassThroughEndpoint client not found in in_memory_llm_clients_cache; " "get_async_httpx_client may not be caching this provider." ) - cache_dict[cache_key] = SimpleNamespace(client=httpx.AsyncClient(transport=httpx.MockTransport(transport_handler))) + cache_dict[cache_key] = SimpleNamespace( + client=httpx.AsyncClient(transport=httpx.MockTransport(transport_handler)) + ) mock_request = MagicMock(spec=Request) mock_request.method = "GET" @@ -2174,14 +2306,18 @@ async def _run_pass_through_and_capture_wire_url( mock_request.body = AsyncMock(return_value=b"") mock_proxy_logging = MagicMock() - mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type: data) + mock_proxy_logging.pre_call_hook = AsyncMock( + side_effect=lambda user_api_key_dict, data, call_type: data + ) mock_proxy_logging.post_call_failure_hook = AsyncMock() mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value={}) mock_proxy_logging.get_proxy_hook = MagicMock(return_value=managed_files_hook) try: with ExitStack() as stack: - stack.enter_context(patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging)) + stack.enter_context( + patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging) + ) if managed_files_hook is not None: stack.enter_context( patch( @@ -2349,16 +2485,26 @@ async def test_filter_endpoints_by_team_allowed_routes_with_filter(): # Create test endpoints endpoints = [ - PassThroughGenericEndpoint(id="endpoint-1", path="/api/allowed1", target="http://example.com/api1"), - PassThroughGenericEndpoint(id="endpoint-2", path="/api/allowed2", target="http://example.com/api2"), - PassThroughGenericEndpoint(id="endpoint-3", path="/api/notallowed", target="http://example.com/api3"), + PassThroughGenericEndpoint( + id="endpoint-1", path="/api/allowed1", target="http://example.com/api1" + ), + PassThroughGenericEndpoint( + id="endpoint-2", path="/api/allowed2", target="http://example.com/api2" + ), + PassThroughGenericEndpoint( + id="endpoint-3", path="/api/notallowed", target="http://example.com/api3" + ), ] # Mock prisma client mock_prisma_client = MagicMock() mock_team = MagicMock() - mock_team.metadata = {"allowed_passthrough_routes": ["/api/allowed1", "/api/allowed2"]} - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_team) + mock_team.metadata = { + "allowed_passthrough_routes": ["/api/allowed1", "/api/allowed2"] + } + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=mock_team + ) # Call the function result = await _filter_endpoints_by_team_allowed_routes( @@ -2373,7 +2519,9 @@ async def test_filter_endpoints_by_team_allowed_routes_with_filter(): assert result[1].path == "/api/allowed2" # Verify database call - mock_prisma_client.db.litellm_teamtable.find_unique.assert_called_once_with(where={"team_id": "test-team-123"}) + mock_prisma_client.db.litellm_teamtable.find_unique.assert_called_once_with( + where={"team_id": "test-team-123"} + ) @pytest.mark.asyncio @@ -2391,7 +2539,9 @@ async def test_filter_endpoints_by_team_allowed_routes_team_not_found(): # Create test endpoints endpoints = [ - PassThroughGenericEndpoint(id="endpoint-1", path="/api/test", target="http://example.com/api"), + PassThroughGenericEndpoint( + id="endpoint-1", path="/api/test", target="http://example.com/api" + ), ] # Mock prisma client to return None (team not found) @@ -2424,15 +2574,21 @@ async def test_filter_endpoints_by_team_allowed_routes_no_metadata(): # Create test endpoints endpoints = [ - PassThroughGenericEndpoint(id="endpoint-1", path="/api/test1", target="http://example.com/api1"), - PassThroughGenericEndpoint(id="endpoint-2", path="/api/test2", target="http://example.com/api2"), + PassThroughGenericEndpoint( + id="endpoint-1", path="/api/test1", target="http://example.com/api1" + ), + PassThroughGenericEndpoint( + id="endpoint-2", path="/api/test2", target="http://example.com/api2" + ), ] # Mock prisma client with team that has None metadata mock_prisma_client = MagicMock() mock_team = MagicMock() mock_team.metadata = None - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_team) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=mock_team + ) # Call the function result = await _filter_endpoints_by_team_allowed_routes( @@ -2460,15 +2616,21 @@ async def test_filter_endpoints_by_team_allowed_routes_no_allowed_routes_key(): # Create test endpoints endpoints = [ - PassThroughGenericEndpoint(id="endpoint-1", path="/api/test1", target="http://example.com/api1"), - PassThroughGenericEndpoint(id="endpoint-2", path="/api/test2", target="http://example.com/api2"), + PassThroughGenericEndpoint( + id="endpoint-1", path="/api/test1", target="http://example.com/api1" + ), + PassThroughGenericEndpoint( + id="endpoint-2", path="/api/test2", target="http://example.com/api2" + ), ] # Mock prisma client with team that has metadata but no allowed_passthrough_routes mock_prisma_client = MagicMock() mock_team = MagicMock() mock_team.metadata = {"some_other_key": "some_value"} - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_team) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=mock_team + ) # Call the function result = await _filter_endpoints_by_team_allowed_routes( @@ -2496,15 +2658,21 @@ async def test_filter_endpoints_by_team_allowed_routes_empty_allowed_list(): # Create test endpoints endpoints = [ - PassThroughGenericEndpoint(id="endpoint-1", path="/api/test1", target="http://example.com/api1"), - PassThroughGenericEndpoint(id="endpoint-2", path="/api/test2", target="http://example.com/api2"), + PassThroughGenericEndpoint( + id="endpoint-1", path="/api/test1", target="http://example.com/api1" + ), + PassThroughGenericEndpoint( + id="endpoint-2", path="/api/test2", target="http://example.com/api2" + ), ] # Mock prisma client with team that has empty allowed_passthrough_routes mock_prisma_client = MagicMock() mock_team = MagicMock() mock_team.metadata = {"allowed_passthrough_routes": []} - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_team) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=mock_team + ) # Call the function result = await _filter_endpoints_by_team_allowed_routes( @@ -2530,21 +2698,29 @@ async def test_filter_endpoints_by_team_allowed_routes_partial_match(): # Create test endpoints endpoints = [ - PassThroughGenericEndpoint(id="endpoint-1", path="/api/openai", target="http://example.com/openai"), + PassThroughGenericEndpoint( + id="endpoint-1", path="/api/openai", target="http://example.com/openai" + ), PassThroughGenericEndpoint( id="endpoint-2", path="/api/anthropic", target="http://example.com/anthropic", ), - PassThroughGenericEndpoint(id="endpoint-3", path="/api/azure", target="http://example.com/azure"), - PassThroughGenericEndpoint(id="endpoint-4", path="/api/cohere", target="http://example.com/cohere"), + PassThroughGenericEndpoint( + id="endpoint-3", path="/api/azure", target="http://example.com/azure" + ), + PassThroughGenericEndpoint( + id="endpoint-4", path="/api/cohere", target="http://example.com/cohere" + ), ] # Mock prisma client with team that allows only 2 routes mock_prisma_client = MagicMock() mock_team = MagicMock() mock_team.metadata = {"allowed_passthrough_routes": ["/api/openai", "/api/azure"]} - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_team) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=mock_team + ) # Call the function result = await _filter_endpoints_by_team_allowed_routes( @@ -2576,7 +2752,9 @@ async def test_bedrock_router_passthrough_metadata_initialization(): ) # Mock ProxyBaseLLMRequestProcessing to verify it's used - with patch("litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing") as mock_processing_class: + with patch( + "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing" + ) as mock_processing_class: # Setup mock instance mock_processor = MagicMock() mock_processing_class.return_value = mock_processor @@ -2584,8 +2762,12 @@ async def test_bedrock_router_passthrough_metadata_initialization(): # Mock successful response mock_response = MagicMock() mock_response.status_code = 200 - mock_response.aread = AsyncMock(return_value=b'{"content": [{"text": "Hello"}]}') - mock_processor.base_passthrough_process_llm_request = AsyncMock(return_value=mock_response) + mock_response.aread = AsyncMock( + return_value=b'{"content": [{"text": "Hello"}]}' + ) + mock_processor.base_passthrough_process_llm_request = AsyncMock( + return_value=mock_response + ) # Create mock request with headers mock_request = MagicMock(spec=Request) @@ -2652,10 +2834,18 @@ async def test_bedrock_router_passthrough_metadata_initialization(): call_kwargs = mock_processor.base_passthrough_process_llm_request.call_args[1] # These are the critical parameters that ensure metadata is properly initialized: - assert call_kwargs["request"] == mock_request, "Request must be passed for header extraction" - assert call_kwargs["user_api_key_dict"] == mock_user_api_key_dict, "User API key dict needed for metadata" - assert call_kwargs["proxy_logging_obj"] == mock_proxy_logging, "Logging obj needed for hooks" - assert call_kwargs["llm_router"] == mock_router, "Router needed for model routing" + assert ( + call_kwargs["request"] == mock_request + ), "Request must be passed for header extraction" + assert ( + call_kwargs["user_api_key_dict"] == mock_user_api_key_dict + ), "User API key dict needed for metadata" + assert ( + call_kwargs["proxy_logging_obj"] == mock_proxy_logging + ), "Logging obj needed for hooks" + assert ( + call_kwargs["llm_router"] == mock_router + ), "Router needed for model routing" assert call_kwargs["model"] == "my-bedrock-model", "Model name must be passed" # Verify response was returned @@ -2718,12 +2908,18 @@ async def test_add_litellm_data_to_request_adds_headers_to_metadata(): # Bedrock passthrough uses litellm_metadata to prevent key-level # tags from leaking into the provider payload (GH#30629). assert "litellm_metadata" in result, "litellm_metadata should be present in result" - assert "headers" in result["litellm_metadata"], "headers should be present in litellm_metadata" - assert isinstance(result["litellm_metadata"]["headers"], dict), "headers should be a dictionary" + assert ( + "headers" in result["litellm_metadata"] + ), "headers should be present in litellm_metadata" + assert isinstance( + result["litellm_metadata"]["headers"], dict + ), "headers should be a dictionary" # Verify specific headers are accessible (important for guardrails) headers = result["litellm_metadata"]["headers"] - assert "user-agent" in headers or "User-Agent" in headers, "User-Agent header should be accessible in metadata" + assert ( + "user-agent" in headers or "User-Agent" in headers + ), "User-Agent header should be accessible in metadata" # Also verify proxy_server_request has headers (original location) assert "proxy_server_request" in result @@ -2758,7 +2954,9 @@ async def test_create_pass_through_route_custom_body_url_target(): ) with ( - patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_request") as mock_pass_through, + patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_request" + ) as mock_pass_through, patch( "litellm.proxy.pass_through_endpoints.pass_through_endpoints.InitPassThroughEndpointHelpers.is_registered_pass_through_route" ) as mock_is_registered, @@ -2795,7 +2993,9 @@ async def test_create_pass_through_route_custom_body_url_target(): "retrievalQuery": {"text": "What is in the knowledge base?"}, } - setattr(mock_request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, bedrock_body) + setattr( + mock_request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, bedrock_body + ) await endpoint_func( request=mock_request, @@ -2833,7 +3033,9 @@ async def test_pass_through_request_non_streaming_uses_content_for_state_raw_bod mock_request.headers = Headers({"Content-Type": "application/json"}) mock_request.state = SimpleNamespace() setattr(mock_request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, raw_signed) - mock_request.body = AsyncMock(return_value=json.dumps(parsed_from_wire).encode("utf-8")) + mock_request.body = AsyncMock( + return_value=json.dumps(parsed_from_wire).encode("utf-8") + ) mock_user = MagicMock() mock_user.api_key = "sk-test" @@ -2906,7 +3108,9 @@ async def test_pass_through_request_streaming_uses_content_for_state_raw_body(): mock_request.headers = Headers({"Content-Type": "application/json"}) mock_request.state = SimpleNamespace() setattr(mock_request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, raw_signed) - mock_request.body = AsyncMock(return_value=json.dumps(parsed_from_wire).encode("utf-8")) + mock_request.body = AsyncMock( + return_value=json.dumps(parsed_from_wire).encode("utf-8") + ) mock_user = MagicMock() mock_user.api_key = "sk-test" @@ -2980,7 +3184,9 @@ async def test_create_pass_through_route_no_custom_body_falls_back(): ) with ( - patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_request") as mock_pass_through, + patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_request" + ) as mock_pass_through, patch( "litellm.proxy.pass_through_endpoints.pass_through_endpoints.InitPassThroughEndpointHelpers.is_registered_pass_through_route" ) as mock_is_registered, @@ -3049,12 +3255,32 @@ def test_is_registered_pass_through_route_with_custom_root(): } with patch("litellm.proxy.utils.get_server_root_path", return_value="/proxy"): - assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/proxy/api/endpoint") is True - assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/api/endpoint") is True + assert ( + InitPassThroughEndpointHelpers.is_registered_pass_through_route( + "/proxy/api/endpoint" + ) + is True + ) + assert ( + InitPassThroughEndpointHelpers.is_registered_pass_through_route( + "/api/endpoint" + ) + is True + ) with patch("litellm.proxy.utils.get_server_root_path", return_value="/"): - assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/api/endpoint") is True - assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/proxy/api/endpoint") is False + assert ( + InitPassThroughEndpointHelpers.is_registered_pass_through_route( + "/api/endpoint" + ) + is True + ) + assert ( + InitPassThroughEndpointHelpers.is_registered_pass_through_route( + "/proxy/api/endpoint" + ) + is False + ) # Clean up _registered_pass_through_routes.clear() @@ -3086,18 +3312,24 @@ def test_get_registered_pass_through_route_with_custom_root(): with patch("litellm.proxy.utils.get_server_root_path", return_value="/litellm"): # Prefixed incoming route - result = InitPassThroughEndpointHelpers.get_registered_pass_through_route("/litellm/chat/completions") + result = InitPassThroughEndpointHelpers.get_registered_pass_through_route( + "/litellm/chat/completions" + ) assert result is not None assert result["target"] == "http://api.example.com/v1/chat/completions" assert result["headers"]["Authorization"] == "Bearer token123" # Bare incoming route (get_request_route convention) - result = InitPassThroughEndpointHelpers.get_registered_pass_through_route("/chat/completions") + result = InitPassThroughEndpointHelpers.get_registered_pass_through_route( + "/chat/completions" + ) assert result is not None assert result["target"] == "http://api.example.com/v1/chat/completions" with patch("litellm.proxy.utils.get_server_root_path", return_value="/"): - result = InitPassThroughEndpointHelpers.get_registered_pass_through_route("/chat/completions") + result = InitPassThroughEndpointHelpers.get_registered_pass_through_route( + "/chat/completions" + ) assert result is not None assert result["target"] == "http://api.example.com/v1/chat/completions" @@ -3151,7 +3383,12 @@ def test_db_registered_pass_through_route_bare_path_convention( "litellm.proxy.utils.get_server_root_path", return_value=server_root_path, ): - assert InitPassThroughEndpointHelpers.is_registered_pass_through_route(incoming_route) is should_match + assert ( + InitPassThroughEndpointHelpers.is_registered_pass_through_route( + incoming_route + ) + is should_match + ) _registered_pass_through_routes.clear() @@ -3170,13 +3407,25 @@ def test_mapped_pass_through_routes_with_server_root_path(): with patch("litellm.proxy.utils.get_server_root_path", return_value="/litellm"): # prefixed route should match mapped routes like /vertex_ai assert ( - InitPassThroughEndpointHelpers.is_registered_pass_through_route("/litellm/vertex_ai/v1/projects/foo") + InitPassThroughEndpointHelpers.is_registered_pass_through_route( + "/litellm/vertex_ai/v1/projects/foo" + ) + is True + ) + assert ( + InitPassThroughEndpointHelpers.is_registered_pass_through_route( + "/litellm/bedrock/model/invoke" + ) is True ) - assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/litellm/bedrock/model/invoke") is True # bare route without prefix should not match when root is set - assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/vertex_ai/v1/projects/foo") is False + assert ( + InitPassThroughEndpointHelpers.is_registered_pass_through_route( + "/vertex_ai/v1/projects/foo" + ) + is False + ) @pytest.mark.asyncio @@ -3193,18 +3442,24 @@ async def test_multipart_passthrough_preserves_boundary(): mock_response = MagicMock() mock_response.status_code = 200 mock_response.headers = httpx.Headers({"content-type": "application/json"}) - mock_response.aread = AsyncMock(return_value=b'{"filename": "test.txt", "size": 17}') + mock_response.aread = AsyncMock( + return_value=b'{"filename": "test.txt", "size": 17}' + ) mock_response.text = '{"filename": "test.txt", "size": 17}' async def mock_httpx_request(method, url, **kwargs): # Verify that files parameter is passed (not json) assert "files" in kwargs, "Files should be passed for multipart requests" - file_parts = [value for name, value in kwargs["files"] if name == "file"] + file_parts = [ + value for name, value in kwargs["files"] if name == "file" + ] assert len(file_parts) == 1, "File field should be in files" # Verify content-type is NOT in headers (httpx will set it with correct boundary) headers = kwargs.get("headers", {}) - assert "content-type" not in headers, "content-type should be removed for multipart" + assert ( + "content-type" not in headers + ), "content-type should be removed for multipart" filename, content, content_type = file_parts[0] assert filename == "test.txt" @@ -3277,7 +3532,9 @@ def test_get_response_headers_strips_server_and_date(): "connection", "keep-alive", ): - assert stripped not in lowered_keys, f"{stripped!r} must not be forwarded by passthrough" + assert ( + stripped not in lowered_keys + ), f"{stripped!r} must not be forwarded by passthrough" # Application/business headers must still pass through. lowered = {k.lower(): v for k, v in result.items()} @@ -3315,7 +3572,9 @@ class TestStaleRouteCleanupOnReload: ) stack.enter_context(patch("litellm.proxy.proxy_server.premium_user", True)) mock_set_env = stack.enter_context( - patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.set_env_variables_in_header") + patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.set_env_variables_in_header" + ) ) mock_set_env.return_value = {} return stack @@ -3360,10 +3619,14 @@ class TestStaleRouteCleanupOnReload: so the registry would hold both paths instead of only ``/b``. """ with self._patches(): - await initialize_pass_through_endpoints([{"path": "/a", "target": "http://example.com"}]) + await initialize_pass_through_endpoints( + [{"path": "/a", "target": "http://example.com"}] + ) assert self._paths_in_registry() == ["/a"] - await initialize_pass_through_endpoints([{"path": "/b", "target": "http://example.com"}]) + await initialize_pass_through_endpoints( + [{"path": "/b", "target": "http://example.com"}] + ) assert self._paths_in_registry() == ["/b"] @@ -3386,8 +3649,12 @@ class TestStaleRouteCleanupOnReload: ] ) - assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/live-passthrough") - assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/live-passthrough/some/subpath") + assert InitPassThroughEndpointHelpers.is_registered_pass_through_route( + "/live-passthrough" + ) + assert InitPassThroughEndpointHelpers.is_registered_pass_through_route( + "/live-passthrough/some/subpath" + ) # Regression (LIT-3538): a pre-call guardrail block on a passthrough endpoint @@ -3488,26 +3755,40 @@ async def _drive_pass_through_block(raised_exception): 400, ), ( - _FastAPIHTTPException(status_code=400, detail={"error": "Violated moderation policy"}), + _FastAPIHTTPException( + status_code=400, detail={"error": "Violated moderation policy"} + ), 400, ), ], ) -async def test_pre_call_guardrail_block_logs_warning_not_exception(guardrail_exception, expected_code): +async def test_pre_call_guardrail_block_logs_warning_not_exception( + guardrail_exception, expected_code +): status_code, logger = await _drive_pass_through_block(guardrail_exception) assert int(status_code) == expected_code - assert logger.exception.call_count == 0, "guardrail block must not be logged as an ERROR with a traceback" - assert logger.warning.call_count == 1, "guardrail block must be logged once at WARNING" + assert ( + logger.exception.call_count == 0 + ), "guardrail block must not be logged as an ERROR with a traceback" + assert ( + logger.warning.call_count == 1 + ), "guardrail block must be logged once at WARNING" @pytest.mark.asyncio async def test_non_guardrail_exception_still_logs_with_traceback(): - status_code, logger = await _drive_pass_through_block(RuntimeError("upstream connection reset")) + status_code, logger = await _drive_pass_through_block( + RuntimeError("upstream connection reset") + ) assert int(status_code) == 500 - assert logger.exception.call_count == 1, "a genuine failure must still be logged via verbose_proxy_logger.exception" - assert logger.warning.call_count == 0, "a genuine failure must not be downgraded to WARNING" + assert ( + logger.exception.call_count == 1 + ), "a genuine failure must still be logged via verbose_proxy_logger.exception" + assert ( + logger.warning.call_count == 0 + ), "a genuine failure must not be downgraded to WARNING" # Regression: generic config-based passthrough (`pass_through_request`) used to @@ -3546,7 +3827,9 @@ async def test_pass_through_request_non_streaming_upstream_error_returned_unchan ) as mock_success_handler: mock_proxy_logging.pre_call_hook = AsyncMock(return_value={}) mock_proxy_logging.post_call_failure_hook = AsyncMock() - mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None) + mock_proxy_logging.post_call_response_headers_hook = AsyncMock( + return_value=None + ) mock_processing.get_custom_headers.return_value = {} mock_success_handler.return_value = None @@ -3632,7 +3915,9 @@ async def test_pass_through_request_upstream_error_failure_hook_exception_is_swa mock_proxy_logging.post_call_failure_hook = AsyncMock( side_effect=RuntimeError("alerting integration misconfigured") ) - mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None) + mock_proxy_logging.post_call_response_headers_hook = AsyncMock( + return_value=None + ) mock_processing.get_custom_headers.return_value = {} mock_success_handler.return_value = None @@ -3682,7 +3967,9 @@ async def test_pass_through_request_streaming_upstream_error_returned_unchanged( ) as mock_success_handler: mock_proxy_logging.pre_call_hook = AsyncMock(return_value={}) mock_proxy_logging.post_call_failure_hook = AsyncMock() - mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None) + mock_proxy_logging.post_call_response_headers_hook = AsyncMock( + return_value=None + ) mock_success_handler.return_value = None async_client = MagicMock() @@ -3710,7 +3997,10 @@ async def test_pass_through_request_streaming_upstream_error_returned_unchanged( streamed_chunks = [chunk async for chunk in response.body_iterator] await asyncio.sleep(0) - streamed_bytes = b"".join(chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") for chunk in streamed_chunks) + streamed_bytes = b"".join( + chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") + for chunk in streamed_chunks + ) assert streamed_bytes == upstream_content assert json.loads(streamed_bytes) == _UPSTREAM_ERROR_BODY @@ -3756,7 +4046,9 @@ async def test_pass_through_request_non_streaming_success_unchanged(): ) as mock_success_handler: mock_proxy_logging.pre_call_hook = AsyncMock(return_value={}) mock_proxy_logging.post_call_failure_hook = AsyncMock() - mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None) + mock_proxy_logging.post_call_response_headers_hook = AsyncMock( + return_value=None + ) mock_processing.get_custom_headers.return_value = {} mock_success_handler.return_value = None @@ -3800,7 +4092,9 @@ async def test_pass_through_request_internal_failure_still_raises_proxy_exceptio from litellm.proxy._types import ProxyException with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: - mock_proxy_logging.pre_call_hook = AsyncMock(side_effect=RuntimeError("auth backend unavailable")) + mock_proxy_logging.pre_call_hook = AsyncMock( + side_effect=RuntimeError("auth backend unavailable") + ) mock_proxy_logging.post_call_failure_hook = AsyncMock() mock_request = MagicMock(spec=Request) @@ -3889,7 +4183,9 @@ def _inject_fake_passthrough_client(transport, timeout): def _enter_relay_logging_mocks(stack, parsed_body): from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER - mock_proxy_logging = stack.enter_context(patch("litellm.proxy.proxy_server.proxy_logging_obj")) + mock_proxy_logging = stack.enter_context( + patch("litellm.proxy.proxy_server.proxy_logging_obj") + ) mock_proxy_logging.pre_call_hook = AsyncMock(return_value=parsed_body) mock_proxy_logging.post_call_failure_hook = AsyncMock() mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None) @@ -3899,7 +4195,11 @@ def _enter_relay_logging_mocks(stack, parsed_body): ) ) mock_success_handler.return_value = None - stack.enter_context(patch.object(GLOBAL_LOGGING_WORKER, "ensure_initialized_and_enqueue", new=MagicMock())) + stack.enter_context( + patch.object( + GLOBAL_LOGGING_WORKER, "ensure_initialized_and_enqueue", new=MagicMock() + ) + ) return mock_proxy_logging, mock_success_handler @@ -3985,7 +4285,10 @@ async def test_pass_through_request_relays_non_json_body_without_buffering(): mock_success_handler.assert_called_once() success_kwargs = mock_success_handler.call_args.kwargs assert success_kwargs["response_body"] is None - assert success_kwargs["url_route"] == "http://upstream.test/v1/messages/batches/b1/results" + assert ( + success_kwargs["url_route"] + == "http://upstream.test/v1/messages/batches/b1/results" + ) finally: cleanup() await fake_client.aclose() @@ -4063,7 +4366,9 @@ async def test_pass_through_request_upstream_error_body_stays_buffered(): ) try: with ExitStack() as stack: - mock_proxy_logging, mock_success_handler = _enter_relay_logging_mocks(stack, {}) + mock_proxy_logging, mock_success_handler = _enter_relay_logging_mocks( + stack, {} + ) response = await pass_through_request( request=_relay_client_request(), @@ -4133,11 +4438,18 @@ async def test_pass_through_relay_client_disconnect_logs_partial_relay_warning(c partial_relay_warnings = [ record.getMessage() for record in caplog.records - if record.levelno == logging.WARNING and _PARTIAL_RELAY_WARNING_MARKER in record.getMessage() + if record.levelno == logging.WARNING + and _PARTIAL_RELAY_WARNING_MARKER in record.getMessage() ] assert len(partial_relay_warnings) == 1 - assert "http://upstream.test/v1/messages/batches/b1/results" in partial_relay_warnings[0] - assert f"{len(first_chunk)} bytes were sent to the client" in partial_relay_warnings[0] + assert ( + "http://upstream.test/v1/messages/batches/b1/results" + in partial_relay_warnings[0] + ) + assert ( + f"{len(first_chunk)} bytes were sent to the client" + in partial_relay_warnings[0] + ) assert upstream_stream.closed is True mock_success_handler.assert_called_once() @@ -4184,7 +4496,10 @@ async def test_pass_through_relay_full_consumption_logs_no_partial_relay_warning relayed = [chunk async for chunk in response.body_iterator] assert b"".join(relayed) == b"".join(upstream_chunks) - assert not any(_PARTIAL_RELAY_WARNING_MARKER in record.getMessage() for record in caplog.records) + assert not any( + _PARTIAL_RELAY_WARNING_MARKER in record.getMessage() + for record in caplog.records + ) mock_success_handler.assert_called_once() finally: cleanup()