diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index ddf6a59bea7..4a5f66f7f25 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -15705,6 +15705,228 @@ } } }, + "/typesafe/{endpoint}": { + "delete": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/typesafe)", + "operationId": "typesafe_proxy_route_typesafe__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Typesafe Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/typesafe)", + "operationId": "typesafe_proxy_route_typesafe__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Typesafe Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/typesafe)", + "operationId": "typesafe_proxy_route_typesafe__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Typesafe Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/typesafe)", + "operationId": "typesafe_proxy_route_typesafe__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Typesafe Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/typesafe)", + "operationId": "typesafe_proxy_route_typesafe__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Typesafe Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, "mcp_app": { "components": { "schemas": { diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 920900730d1..44a19137234 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -503,7 +503,7 @@ async def mistral_proxy_route( @router.api_route( "/typesafe/{endpoint:path}", - methods=["GET", "POST"], # mutable-ok: FastAPI route metadata requires a list + methods=["GET", "POST", "PUT", "DELETE", "PATCH"], # mutable-ok: FastAPI route metadata requires a list tags=["TypeSafe AI Pass-through", "pass-through"], # mutable-ok: FastAPI route metadata requires a list ) async def typesafe_proxy_route( diff --git a/tests/litellm_utils_tests/test_utils.py b/tests/litellm_utils_tests/test_utils.py index 0ccfae55290..1c5e399e601 100644 --- a/tests/litellm_utils_tests/test_utils.py +++ b/tests/litellm_utils_tests/test_utils.py @@ -75,9 +75,7 @@ def test_basic_trimming_no_max_tokens_specified(): print("trimmed messages for gpt-4") print(trimmed_messages) # print(get_token_count(messages=trimmed_messages, model="claude-2")) - assert ( - get_token_count(messages=trimmed_messages, model="gpt-4") - ) <= litellm.model_cost["gpt-4"]["max_tokens"] + assert (get_token_count(messages=trimmed_messages, model="gpt-4")) <= litellm.model_cost["gpt-4"]["max_tokens"] # test_basic_trimming_no_max_tokens_specified() @@ -94,9 +92,7 @@ def test_multiple_messages_trimming(): "content": "This is another long message that will also exceed the limit.", }, ] - trimmed_messages = trim_messages( - messages=messages, model="gpt-3.5-turbo", max_tokens=20 - ) + trimmed_messages = trim_messages(messages=messages, model="gpt-3.5-turbo", max_tokens=20) # print(get_token_count(messages=trimmed_messages, model="gpt-3.5-turbo")) assert (get_token_count(messages=trimmed_messages, model="gpt-3.5-turbo")) <= 20 @@ -115,9 +111,7 @@ def test_multiple_messages_no_trimming(): "content": "This is another long message that will also exceed the limit.", }, ] - trimmed_messages = trim_messages( - messages=messages, model="gpt-3.5-turbo", max_tokens=100 - ) + trimmed_messages = trim_messages(messages=messages, model="gpt-3.5-turbo", max_tokens=100) print("Trimmed messages") print(trimmed_messages) assert messages == trimmed_messages @@ -144,9 +138,7 @@ def test_large_trimming_multiple_messages(): def test_large_trimming_single_message(): - messages = [ - {"role": "user", "content": "This is a singlelongwordthatexceedsthelimit."} - ] + messages = [{"role": "user", "content": "This is a singlelongwordthatexceedsthelimit."}] trimmed_messages = trim_messages(messages, max_tokens=5, model="gpt-4-0613") assert (get_token_count(messages=trimmed_messages, model="gpt-4-0613")) <= 5 assert (get_token_count(messages=trimmed_messages, model="gpt-4-0613")) > 0 @@ -277,10 +269,7 @@ def test_trimming_with_model_cost_max_input_tokens(model): }, ] trimmed_messages = trim_messages(messages, model=model) - assert ( - get_token_count(trimmed_messages, model=model) - < litellm.model_cost[model]["max_input_tokens"] - ) + assert get_token_count(trimmed_messages, model=model) < litellm.model_cost[model]["max_input_tokens"] def test_trimming_with_untokenizable_field(caplog: pytest.LogCaptureFixture) -> None: @@ -333,9 +322,7 @@ def test_aget_valid_models(): print(valid_models) # list of openai supported llms on litellm - expected_models = ( - litellm.open_ai_chat_completion_models | litellm.open_ai_text_completion_models - ) + expected_models = litellm.open_ai_chat_completion_models | litellm.open_ai_text_completion_models assert set(valid_models) == set(expected_models) @@ -357,9 +344,7 @@ def test_get_valid_models_with_custom_llm_provider(custom_llm_provider): provider=LlmProviders(custom_llm_provider), ) assert provider_config is not None - valid_models = get_valid_models( - check_provider_endpoint=True, custom_llm_provider=custom_llm_provider - ) + valid_models = get_valid_models(check_provider_endpoint=True, custom_llm_provider=custom_llm_provider) print(valid_models) assert len(valid_models) > 0 assert set(provider_config.get_models()) == set(valid_models) @@ -392,9 +377,7 @@ def test_validate_environment_empty_model(): def test_validate_environment_api_key(): response_obj = validate_environment(model="gpt-5-mini", api_key="sk-my-test-key") - assert ( - response_obj["keys_in_environment"] is True - ), f"Missing keys={response_obj['missing_keys']}" + assert response_obj["keys_in_environment"] is True, f"Missing keys={response_obj['missing_keys']}" def test_validate_environment_api_version(): @@ -404,9 +387,7 @@ def test_validate_environment_api_version(): api_base="https://fake.openai.azure.com/", api_version="2024-02-15", ) - assert ( - response_obj["keys_in_environment"] is True - ), f"Missing keys={response_obj['missing_keys']}" + assert response_obj["keys_in_environment"] is True, f"Missing keys={response_obj['missing_keys']}" def test_validate_environment_api_base_dynamic(): @@ -481,18 +462,14 @@ def test_function_to_dict(): assert function_json["description"] == expected_output["description"] assert function_json["parameters"]["type"] == expected_output["parameters"]["type"] assert ( - function_json["parameters"]["properties"]["location"] - == expected_output["parameters"]["properties"]["location"] + function_json["parameters"]["properties"]["location"] == expected_output["parameters"]["properties"]["location"] ) # the enum can change it can be - which is why we don't assert on unit # {'type': 'string', 'description': 'Temperature unit', 'enum': "['fahrenheit', 'celsius']"} # {'type': 'string', 'description': 'Temperature unit', 'enum': "['celsius', 'fahrenheit']"} - assert ( - function_json["parameters"]["required"] - == expected_output["parameters"]["required"] - ) + assert function_json["parameters"]["required"] == expected_output["parameters"]["required"] print("passed") @@ -561,9 +538,7 @@ def test_get_max_token_unit_test(): """ model = "bedrock/anthropic.claude-3-haiku-20240307-v1:0" - max_tokens = get_max_tokens( - model - ) # Returns a number instead of throwing an Exception + max_tokens = get_max_tokens(model) # Returns a number instead of throwing an Exception assert isinstance(max_tokens, int) @@ -602,9 +577,7 @@ def test_get_chat_completion_prompt(): prompt_variables=None, ) - assert litellm_logging_obj.messages == [ - {"role": "user", "content": updated_message} - ] + assert litellm_logging_obj.messages == [{"role": "user", "content": updated_message}] def test_redact_msgs_from_logs(): @@ -676,9 +649,7 @@ def test_redact_embedding_response(): litellm.turn_off_message_logging = True # Create a test EmbeddingResponse with usage data - original_usage = litellm.Usage( - prompt_tokens=10, completion_tokens=0, total_tokens=10 - ) + original_usage = litellm.Usage(prompt_tokens=10, completion_tokens=0, total_tokens=10) original_data = [ {"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3, 0.4, 0.5]}, {"object": "embedding", "index": 1, "embedding": [0.6, 0.7, 0.8, 0.9, 1.0]}, @@ -714,9 +685,7 @@ def test_redact_embedding_response(): # Assert the redacted response preserves critical metadata assert _redacted_response_obj.usage == original_usage # usage should be preserved - assert ( - _redacted_response_obj.model == "text-embedding-3-small" - ) # model should be preserved + assert _redacted_response_obj.model == "text-embedding-3-small" # model should be preserved assert _redacted_response_obj.object == "list" # object should be preserved # Assert sensitive data is cleared @@ -770,12 +739,8 @@ def test_redact_msgs_from_logs_with_dynamic_params(): ) # Test Case 1: standard_callback_dynamic_params = False (or not set) - standard_callback_dynamic_params = StandardCallbackDynamicParams( - turn_off_message_logging=False - ) - litellm_logging_obj.model_call_details["standard_callback_dynamic_params"] = ( - standard_callback_dynamic_params - ) + standard_callback_dynamic_params = StandardCallbackDynamicParams(turn_off_message_logging=False) + litellm_logging_obj.model_call_details["standard_callback_dynamic_params"] = standard_callback_dynamic_params _redacted_response_obj = redact_message_input_output_from_logging( result=response_obj, model_call_details=litellm_logging_obj.model_call_details, @@ -784,12 +749,8 @@ def test_redact_msgs_from_logs_with_dynamic_params(): assert _redacted_response_obj.choices[0].message.content == test_content # Test Case 2: standard_callback_dynamic_params = True - standard_callback_dynamic_params = StandardCallbackDynamicParams( - turn_off_message_logging=True - ) - litellm_logging_obj.model_call_details["standard_callback_dynamic_params"] = ( - standard_callback_dynamic_params - ) + standard_callback_dynamic_params = StandardCallbackDynamicParams(turn_off_message_logging=True) + litellm_logging_obj.model_call_details["standard_callback_dynamic_params"] = standard_callback_dynamic_params _redacted_response_obj = redact_message_input_output_from_logging( result=response_obj, model_call_details=litellm_logging_obj.model_call_details, @@ -800,9 +761,7 @@ def test_redact_msgs_from_logs_with_dynamic_params(): # Test Case 3: standard_callback_dynamic_params does not set turn_off_message_logging # since litellm.turn_off_message_logging is True redaction should occur standard_callback_dynamic_params = StandardCallbackDynamicParams() - litellm_logging_obj.model_call_details["standard_callback_dynamic_params"] = ( - standard_callback_dynamic_params - ) + litellm_logging_obj.model_call_details["standard_callback_dynamic_params"] = standard_callback_dynamic_params _redacted_response_obj = redact_message_input_output_from_logging( result=response_obj, model_call_details=litellm_logging_obj.model_call_details, @@ -907,9 +866,7 @@ def test_get_llm_provider_ft_models(): @pytest.mark.parametrize("langfuse_trace_id", [None, "my-unique-trace-id"]) -@pytest.mark.parametrize( - "langfuse_existing_trace_id", [None, "my-unique-existing-trace-id"] -) +@pytest.mark.parametrize("langfuse_existing_trace_id", [None, "my-unique-existing-trace-id"]) def test_logging_trace_id(langfuse_trace_id, langfuse_existing_trace_id): """ - Unit test for `_get_trace_id` function in Logging obj @@ -948,22 +905,13 @@ def test_logging_trace_id(langfuse_trace_id, langfuse_existing_trace_id): ## if existing_trace_id exists if langfuse_existing_trace_id is not None: - assert ( - litellm_logging_obj._get_trace_id(service_name="langfuse") - == langfuse_existing_trace_id - ) + assert litellm_logging_obj._get_trace_id(service_name="langfuse") == langfuse_existing_trace_id ## if trace_id exists elif langfuse_trace_id is not None: - assert ( - litellm_logging_obj._get_trace_id(service_name="langfuse") - == langfuse_trace_id - ) + assert litellm_logging_obj._get_trace_id(service_name="langfuse") == langfuse_trace_id ## if no trace_id or existing_trace_id is provided, use litellm_trace_id else: - assert ( - litellm_logging_obj._get_trace_id(service_name="langfuse") - == litellm_logging_obj.litellm_trace_id - ) + assert litellm_logging_obj._get_trace_id(service_name="langfuse") == litellm_logging_obj.litellm_trace_id def test_convert_model_response_object(): @@ -1154,9 +1102,7 @@ def test_async_http_handler(mock_async_client): concurrent_limit = 2 # Mock the transport creation to return a specific transport - with mock.patch.object( - AsyncHTTPHandler, "_create_async_transport" - ) as mock_create_transport: + with mock.patch.object(AsyncHTTPHandler, "_create_async_transport") as mock_create_transport: mock_transport = mock.MagicMock() mock_create_transport.return_value = mock_transport @@ -1221,9 +1167,7 @@ def test_async_http_handler_force_ipv4(mock_async_client): litellm.force_ipv4 = False -@pytest.mark.parametrize( - "model, expected_bool", [("gpt-3.5-turbo", False), ("gpt-4o-audio-preview", True)] -) +@pytest.mark.parametrize("model, expected_bool", [("gpt-3.5-turbo", False), ("gpt-4o-audio-preview", True)]) def test_supports_audio_input(model, expected_bool): os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") @@ -1277,9 +1221,7 @@ def test_is_base64_encoded_2(): [ { "role": "user", - "content": [ - {"type": "image_url", "url": "https://example.com/image.png"} - ], + "content": [{"type": "image_url", "url": "https://example.com/image.png"}], } ], True, @@ -1355,21 +1297,15 @@ def test_models_by_provider(): continue elif k == "sample_spec": continue - elif ( - v["litellm_provider"] == "sagemaker" - or v["litellm_provider"] == "bedrock_converse" - ): + elif v["litellm_provider"] == "sagemaker" or v["litellm_provider"] == "bedrock_converse": continue - elif v.get("mode") == "search": - # Skip search providers as they don't have traditional models + elif v.get("mode") in ("search", "evaluation"): continue else: providers.add(v["litellm_provider"]) for provider in providers: - assert provider in models_by_provider.keys() or JSONProviderRegistry.exists( - provider - ) + assert provider in models_by_provider.keys() or JSONProviderRegistry.exists(provider) @pytest.mark.parametrize( @@ -1380,16 +1316,11 @@ def test_models_by_provider(): ({"user_api_key_end_user_id": "123"}, True, None), ], ) -def test_get_end_user_id_for_cost_tracking( - litellm_params, disable_end_user_cost_tracking, expected_end_user_id -): +def test_get_end_user_id_for_cost_tracking(litellm_params, disable_end_user_cost_tracking, expected_end_user_id): from litellm.utils import get_end_user_id_for_cost_tracking litellm.disable_end_user_cost_tracking = disable_end_user_cost_tracking - assert ( - get_end_user_id_for_cost_tracking(litellm_params=litellm_params) - == expected_end_user_id - ) + assert get_end_user_id_for_cost_tracking(litellm_params=litellm_params) == expected_end_user_id @pytest.mark.parametrize( @@ -1405,13 +1336,9 @@ def test_get_end_user_id_for_cost_tracking_prometheus_only( ): from litellm.utils import get_end_user_id_for_cost_tracking - litellm.enable_end_user_cost_tracking_prometheus_only = ( - enable_end_user_cost_tracking_prometheus_only - ) + litellm.enable_end_user_cost_tracking_prometheus_only = enable_end_user_cost_tracking_prometheus_only assert ( - get_end_user_id_for_cost_tracking( - litellm_params=litellm_params, service_type="prometheus" - ) + get_end_user_id_for_cost_tracking(litellm_params=litellm_params, service_type="prometheus") == expected_end_user_id ) @@ -1426,20 +1353,14 @@ def test_get_end_user_id_for_cost_tracking_prometheus_only( ), # Test with only litellm_metadata field (new behavior) ( - { - "litellm_metadata": { - "user_api_key_end_user_id": "user_from_litellm_metadata" - } - }, + {"litellm_metadata": {"user_api_key_end_user_id": "user_from_litellm_metadata"}}, "user_from_litellm_metadata", ), # Test with both fields - metadata should take precedence for user_api_key fields ( { "metadata": {"user_api_key_end_user_id": "user_from_metadata"}, - "litellm_metadata": { - "user_api_key_end_user_id": "user_from_litellm_metadata" - }, + "litellm_metadata": {"user_api_key_end_user_id": "user_from_litellm_metadata"}, }, "user_from_metadata", ), @@ -1455,9 +1376,7 @@ def test_get_end_user_id_for_cost_tracking_prometheus_only( ( { "metadata": {}, - "litellm_metadata": { - "user_api_key_end_user_id": "user_from_litellm_metadata" - }, + "litellm_metadata": {"user_api_key_end_user_id": "user_from_litellm_metadata"}, }, "user_from_litellm_metadata", ), @@ -1465,9 +1384,7 @@ def test_get_end_user_id_for_cost_tracking_prometheus_only( ({}, None), ], ) -def test_get_end_user_id_for_cost_tracking_metadata_handling( - litellm_params, expected_end_user_id -): +def test_get_end_user_id_for_cost_tracking_metadata_handling(litellm_params, expected_end_user_id): """ Test that get_end_user_id_for_cost_tracking correctly handles both metadata and litellm_metadata fields using the get_litellm_metadata_from_kwargs helper function. @@ -1631,9 +1548,7 @@ def test_get_valid_models_openai_proxy(monkeypatch): mock_response.status_code = 200 mock_response.json.return_value = mock_response_data - with patch.object( - litellm.module_level_client, "get", return_value=mock_response - ) as mock_post: + with patch.object(litellm.module_level_client, "get", return_value=mock_response) as mock_post: valid_models = get_valid_models(check_provider_endpoint=True) assert "litellm_proxy/gpt-5.5" in valid_models @@ -1710,16 +1625,11 @@ def test_get_valid_models_fireworks_ai(monkeypatch): mock_response.status_code = 200 mock_response.json.return_value = mock_response_data - with patch.object( - litellm.module_level_client, "get", return_value=mock_response - ) as mock_post: + with patch.object(litellm.module_level_client, "get", return_value=mock_response) as mock_post: valid_models = get_valid_models(check_provider_endpoint=True) print("valid_models", valid_models) mock_post.assert_called_once() - assert ( - "fireworks_ai/accounts/fireworks/models/llama-3.1-8b-instruct" - in valid_models - ) + assert "fireworks_ai/accounts/fireworks/models/llama-3.1-8b-instruct" in valid_models def test_get_valid_models_default(monkeypatch): @@ -1758,9 +1668,7 @@ def test_pick_cheapest_chat_model_from_llm_provider(): def test_get_num_retries(num_retries): from litellm.utils import _get_wrapper_num_retries - assert _get_wrapper_num_retries( - kwargs={"num_retries": num_retries}, exception=Exception("test") - ) == ( + assert _get_wrapper_num_retries(kwargs={"num_retries": num_retries}, exception=Exception("test")) == ( num_retries, { "num_retries": num_retries, @@ -2033,9 +1941,7 @@ def test_add_custom_logger_callback_to_specific_event_e2e_failure(monkeypatch): assert len(litellm.success_callback) == curr_len_success_callback assert len(litellm.failure_callback) == curr_len_failure_callback - assert any( - isinstance(callback, OpenMeterLogger) for callback in litellm.failure_callback - ) + assert any(isinstance(callback, OpenMeterLogger) for callback in litellm.failure_callback) @pytest.mark.asyncio @@ -2062,20 +1968,13 @@ async def test_wrapper_kwargs_passthrough(): mock_original.assert_called_once() # get litellm logging object - litellm_logging_obj: LiteLLMLoggingObject = mock_original.call_args.kwargs.get( - "litellm_logging_obj" - ) + litellm_logging_obj: LiteLLMLoggingObject = mock_original.call_args.kwargs.get("litellm_logging_obj") assert litellm_logging_obj is not None - print( - f"litellm_logging_obj.model_call_details: {litellm_logging_obj.model_call_details}" - ) + print(f"litellm_logging_obj.model_call_details: {litellm_logging_obj.model_call_details}") # get base model - assert ( - litellm_logging_obj.model_call_details["litellm_params"]["base_model"] - == "gpt-5-mini" - ) + assert litellm_logging_obj.model_call_details["litellm_params"]["base_model"] == "gpt-5-mini" def test_dict_to_response_format_helper(): @@ -2129,7 +2028,7 @@ def test_validate_user_messages_invalid_content_type(): messages = [{"content": [{"type": "invalid_type", "text": "Hello"}]}] - with pytest.raises(Exception, match='Please ensure all messages are valid OpenAI chat completion') as e: + with pytest.raises(Exception, match="Please ensure all messages are valid OpenAI chat completion") as e: validate_chat_completion_user_messages(messages) assert "Invalid message" in str(e) @@ -2146,20 +2045,14 @@ from unittest.mock import Mock [ { "name": "default_on_guardrail", - "callbacks": [ - CustomGuardrail(guardrail_name="test_guardrail", default_on=True) - ], + "callbacks": [CustomGuardrail(guardrail_name="test_guardrail", default_on=True)], "kwargs": {"metadata": {"requester_metadata": {"guardrails": []}}}, "expected": ["test_guardrail"], }, { "name": "request_specific_guardrail", - "callbacks": [ - CustomGuardrail(guardrail_name="test_guardrail", default_on=False) - ], - "kwargs": { - "metadata": {"requester_metadata": {"guardrails": ["test_guardrail"]}} - }, + "callbacks": [CustomGuardrail(guardrail_name="test_guardrail", default_on=False)], + "kwargs": {"metadata": {"requester_metadata": {"guardrails": ["test_guardrail"]}}}, "expected": ["test_guardrail"], }, { @@ -2168,18 +2061,12 @@ from unittest.mock import Mock CustomGuardrail(guardrail_name="default_guardrail", default_on=True), CustomGuardrail(guardrail_name="request_guardrail", default_on=False), ], - "kwargs": { - "metadata": { - "requester_metadata": {"guardrails": ["request_guardrail"]} - } - }, + "kwargs": {"metadata": {"requester_metadata": {"guardrails": ["request_guardrail"]}}}, "expected": ["default_guardrail", "request_guardrail"], }, { "name": "empty_metadata", - "callbacks": [ - CustomGuardrail(guardrail_name="test_guardrail", default_on=False) - ], + "callbacks": [CustomGuardrail(guardrail_name="test_guardrail", default_on=False)], "kwargs": {}, "expected": [], }, @@ -2286,9 +2173,7 @@ def test_get_provider_audio_transcription_config(): from litellm.types.utils import LlmProviders for provider in LlmProviders: - config = ProviderConfigManager.get_provider_audio_transcription_config( - model="whisper-1", provider=provider - ) + config = ProviderConfigManager.get_provider_audio_transcription_config(model="whisper-1", provider=provider) @pytest.mark.parametrize( @@ -2331,9 +2216,7 @@ def test_get_valid_models_from_provider_cache_invalidation(monkeypatch): monkeypatch.setenv("OPENAI_API_KEY", "123") - _model_cache.set_cached_model_info( - "openai", litellm_params=None, available_models=["gpt-5-mini"] - ) + _model_cache.set_cached_model_info("openai", litellm_params=None, available_models=["gpt-5-mini"]) monkeypatch.delenv("OPENAI_API_KEY") assert _model_cache.get_cached_model_info("openai") is None @@ -2422,12 +2305,8 @@ def test_delta_tool_calls_sequential_indices(): # Verify tool calls have sequential indices assert delta.tool_calls is not None, "Tool calls should not be None" assert len(delta.tool_calls) == 2 - assert ( - delta.tool_calls[0].index == 0 - ), f"First tool call should have index 0, got {delta.tool_calls[0].index}" - assert ( - delta.tool_calls[1].index == 1 - ), f"Second tool call should have index 1, got {delta.tool_calls[1].index}" + assert delta.tool_calls[0].index == 0, f"First tool call should have index 0, got {delta.tool_calls[0].index}" + assert delta.tool_calls[1].index == 1, f"Second tool call should have index 1, got {delta.tool_calls[1].index}" # Verify tool call details are preserved assert delta.tool_calls[0].function.name == "get_weather_for_dallas" @@ -2440,9 +2319,7 @@ def test_completion_with_no_model(): """ # test on empty with pytest.raises(TypeError): - response = litellm.completion( - messages=[{"role": "user", "content": "Hello, how are you?"}] - ) + response = litellm.completion(messages=[{"role": "user", "content": "Hello, how are you?"}]) def test_get_base_model_from_metadata(): @@ -2455,43 +2332,31 @@ def test_get_base_model_from_metadata(): from litellm.utils import _get_base_model_from_metadata # Test 1: base_model in metadata (Chat Completions API pattern) - model_call_details_with_metadata = { - "litellm_params": {"metadata": {"model_info": {"base_model": "azure/gpt-5.5"}}} - } + model_call_details_with_metadata = {"litellm_params": {"metadata": {"model_info": {"base_model": "azure/gpt-5.5"}}}} result = _get_base_model_from_metadata(model_call_details_with_metadata) assert result == "azure/gpt-5.5", f"Expected 'azure/gpt-5.5', got {result}" # Test 2: base_model in litellm_metadata (Responses API and generic API calls pattern) model_call_details_with_litellm_metadata = { - "litellm_params": { - "litellm_metadata": {"model_info": {"base_model": "azure/gpt-5-mini"}} - } + "litellm_params": {"litellm_metadata": {"model_info": {"base_model": "azure/gpt-5-mini"}}} } result = _get_base_model_from_metadata(model_call_details_with_litellm_metadata) assert result == "azure/gpt-5-mini", f"Expected 'azure/gpt-5-mini', got {result}" # Test 3: base_model in litellm_params (direct base_model) - model_call_details_with_direct_base_model = { - "litellm_params": {"base_model": "azure/gpt-5-mini"} - } + model_call_details_with_direct_base_model = {"litellm_params": {"base_model": "azure/gpt-5-mini"}} result = _get_base_model_from_metadata(model_call_details_with_direct_base_model) - assert ( - result == "azure/gpt-5-mini" - ), f"Expected 'azure/gpt-5-mini', got {result}" + assert result == "azure/gpt-5-mini", f"Expected 'azure/gpt-5-mini', got {result}" # Test 4: metadata takes precedence over litellm_metadata model_call_details_with_both = { "litellm_params": { "metadata": {"model_info": {"base_model": "azure/gpt-4-from-metadata"}}, - "litellm_metadata": { - "model_info": {"base_model": "azure/gpt-4-from-litellm-metadata"} - }, + "litellm_metadata": {"model_info": {"base_model": "azure/gpt-4-from-litellm-metadata"}}, } } result = _get_base_model_from_metadata(model_call_details_with_both) - assert ( - result == "azure/gpt-4-from-metadata" - ), f"Expected metadata to take precedence, got {result}" + assert result == "azure/gpt-4-from-metadata", f"Expected metadata to take precedence, got {result}" # Test 5: No base_model present model_call_details_without_base_model = {"litellm_params": {"metadata": {}}} diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 84d3aa24871..7c7f6c125ff 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -3,7 +3,7 @@ import contextlib import json import os import traceback -from collections.abc import Mapping +from collections.abc import Iterator, Mapping from types import MappingProxyType, SimpleNamespace from typing import Final from unittest import mock @@ -12,6 +12,7 @@ from urllib.parse import parse_qs import httpx import pytest +import respx from fastapi import HTTPException, Request, Response from fastapi.responses import StreamingResponse from fastapi.testclient import TestClient @@ -45,6 +46,7 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( ) from litellm.proxy._types import LitellmUserRoles, SpecialHeaders, UserAPIKeyAuth from litellm.proxy.auth.handle_jwt import JWTHandler +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials @@ -84,36 +86,28 @@ class TestBaseOpenAIPassThroughHandler: # Test joining base URL with no path and a path base_url = httpx.URL("https://api.example.com") path = "/v1/chat/completions" - result = _join_url_paths( - base_url, path, litellm.LlmProviders.OPENAI.value - ) + result = _join_url_paths(base_url, path, litellm.LlmProviders.OPENAI.value) print(f"Base URL with no path: '{base_url}' + '{path}' → '{result}'") assert str(result) == "https://api.example.com/v1/chat/completions" # Test joining base URL with path and another path base_url = httpx.URL("https://api.example.com/v1") path = "/chat/completions" - result = _join_url_paths( - base_url, path, litellm.LlmProviders.OPENAI.value - ) + result = _join_url_paths(base_url, path, litellm.LlmProviders.OPENAI.value) print(f"Base URL with path: '{base_url}' + '{path}' → '{result}'") assert str(result) == "https://api.example.com/v1/chat/completions" # Test with path not starting with slash base_url = httpx.URL("https://api.example.com/v1") path = "chat/completions" - result = _join_url_paths( - base_url, path, litellm.LlmProviders.OPENAI.value - ) + result = _join_url_paths(base_url, path, litellm.LlmProviders.OPENAI.value) print(f"Path without leading slash: '{base_url}' + '{path}' → '{result}'") assert str(result) == "https://api.example.com/v1/chat/completions" # Test with base URL having trailing slash base_url = httpx.URL("https://api.example.com/v1/") path = "/chat/completions" - result = _join_url_paths( - base_url, path, litellm.LlmProviders.OPENAI.value - ) + result = _join_url_paths(base_url, path, litellm.LlmProviders.OPENAI.value) print(f"Base URL with trailing slash: '{base_url}' + '{path}' → '{result}'") assert str(result) == "https://api.example.com/v1/chat/completions" @@ -132,17 +126,13 @@ class TestBaseOpenAIPassThroughHandler: headers = {"authorization": "Bearer test_key"} # Test with assistants API request - result = BaseOpenAIPassThroughHandler._append_openai_beta_header( - headers, assistants_request - ) + result = BaseOpenAIPassThroughHandler._append_openai_beta_header(headers, assistants_request) print(f"Assistants API request: Added header: {result}") assert result["OpenAI-Beta"] == "assistants=v2" # Test with non-assistants API request headers = {"authorization": "Bearer test_key"} - result = BaseOpenAIPassThroughHandler._append_openai_beta_header( - headers, non_assistants_request - ) + result = BaseOpenAIPassThroughHandler._append_openai_beta_header(headers, non_assistants_request) print(f"Non-assistants API request: Headers: {result}") assert "OpenAI-Beta" not in result @@ -152,9 +142,7 @@ class TestBaseOpenAIPassThroughHandler: assistant_request.url.path = "/v1/assistants/asst_123456" headers = {"authorization": "Bearer test_key"} - result = BaseOpenAIPassThroughHandler._append_openai_beta_header( - headers, assistant_request - ) + result = BaseOpenAIPassThroughHandler._append_openai_beta_header(headers, assistant_request) print(f"Assistant API request: Added header: {result}") assert result["OpenAI-Beta"] == "assistants=v2" @@ -175,9 +163,7 @@ class TestBaseOpenAIPassThroughHandler: "test-header": "value", }, ): - result = BaseOpenAIPassThroughHandler._assemble_headers( - api_key, mock_request - ) + result = BaseOpenAIPassThroughHandler._assemble_headers(api_key, mock_request) print(f"Assembled headers: {result}") assert result["authorization"] == "Bearer test_api_key" assert result["api-key"] == "test_api_key" @@ -217,9 +203,7 @@ class TestBaseOpenAIPassThroughHandler: # Verify create_pass_through_route was called with correct parameters call_args = mock_create_pass_through.call_args[1] - print( - f"create_pass_through_route called with endpoint: {call_args['endpoint']}" - ) + print(f"create_pass_through_route called with endpoint: {call_args['endpoint']}") print(f"create_pass_through_route called with target: {call_args['target']}") assert call_args["endpoint"] == "/chat/completions" assert call_args["target"] == "https://api.openai.com/v1/chat/completions" @@ -271,9 +255,7 @@ class TestVertexAIPassThroughHandler: # Mock request mock_request = Mock() - mock_request.state = ( - None # Prevent Mock from returning a truthy _cached_headers - ) + mock_request.state = None # Prevent Mock from returning a truthy _cached_headers mock_request.method = "POST" mock_request.headers = { "Authorization": "Bearer test-creds", @@ -291,9 +273,7 @@ class TestVertexAIPassThroughHandler: test_token = vertex_credentials with ( - mock.patch( - "litellm.llms.vertex_ai.vertex_llm_base.VertexBase.load_auth" - ) as mock_load_auth, + mock.patch("litellm.llms.vertex_ai.vertex_llm_base.VertexBase.load_auth") as mock_load_auth, mock.patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, @@ -376,9 +356,7 @@ class TestVertexAIPassThroughHandler: # Mock request mock_request = Mock() - mock_request.state = ( - None # Prevent Mock from returning a truthy _cached_headers - ) + mock_request.state = None # Prevent Mock from returning a truthy _cached_headers mock_request.method = "POST" mock_request.headers = { "Authorization": "Bearer test-creds", @@ -396,9 +374,7 @@ class TestVertexAIPassThroughHandler: test_token = vertex_credentials with ( - mock.patch( - "litellm.llms.vertex_ai.vertex_llm_base.VertexBase.load_auth" - ) as mock_load_auth, + mock.patch("litellm.llms.vertex_ai.vertex_llm_base.VertexBase.load_auth") as mock_load_auth, mock.patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, @@ -423,9 +399,7 @@ class TestVertexAIPassThroughHandler: # Mock the vertex handler for global location mock_handler = Mock() - mock_handler.get_default_base_target_url.return_value = ( - "https://aiplatform.googleapis.com/" - ) + mock_handler.get_default_base_target_url.return_value = "https://aiplatform.googleapis.com/" mock_get_handler.return_value = mock_handler # Mock create_pass_through_route to return a function that returns a mock response @@ -459,9 +433,7 @@ class TestVertexAIPassThroughHandler: ], ) @pytest.mark.asyncio - async def test_vertex_passthrough_with_default_credentials( - self, monkeypatch, initial_endpoint - ): + async def test_vertex_passthrough_with_default_credentials(self, monkeypatch, initial_endpoint): """ Test that when no passthrough credentials are set, default credentials are used in the request """ @@ -500,9 +472,7 @@ class TestVertexAIPassThroughHandler: mock_response = Response() with ( - mock.patch( - "litellm.llms.vertex_ai.vertex_llm_base.VertexBase.load_auth" - ) as mock_load_auth, + mock.patch("litellm.llms.vertex_ai.vertex_llm_base.VertexBase.load_auth") as mock_load_auth, mock.patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, @@ -644,17 +614,13 @@ class TestVertexAIPassThroughHandler: mock_request.method = "POST" mock_response = Mock() - with patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.user_api_key_auth" - ) as mock_auth: + with patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.user_api_key_auth") as mock_auth: mock_auth.return_value = {"api_key": "test-key-123"} with patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_pass_through: - mock_pass_through.return_value = AsyncMock( - return_value={"status": "success"} - ) + mock_pass_through.return_value = AsyncMock(return_value={"status": "success"}) with pytest.raises(HTTPException) as exc_info: await vertex_proxy_route( @@ -714,7 +680,9 @@ class TestVertexAIPassThroughHandler: mock_logging_obj.model_call_details = {} # Test URL with multimodal embedding model - url_route = "/v1/projects/test-project/locations/us-central1/publishers/google/models/multimodalembedding@001:predict" + url_route = ( + "/v1/projects/test-project/locations/us-central1/publishers/google/models/multimodalembedding@001:predict" + ) start_time = datetime.datetime.now() end_time = datetime.datetime.now() @@ -732,19 +700,13 @@ class TestVertexAIPassThroughHandler: mock_embedding_response = EmbeddingResponse( object="list", data=[ - Embedding( - embedding=[0.1, 0.2, 0.3, 0.4, 0.5], index=0, object="embedding" - ), - Embedding( - embedding=[0.6, 0.7, 0.8, 0.9, 1.0], index=1, object="embedding" - ), + Embedding(embedding=[0.1, 0.2, 0.3, 0.4, 0.5], index=0, object="embedding"), + Embedding(embedding=[0.6, 0.7, 0.8, 0.9, 1.0], index=1, object="embedding"), ], model="multimodalembedding@001", usage=Usage(prompt_tokens=0, total_tokens=0, completion_tokens=0), ) - mock_config_instance.transform_embedding_response.return_value = ( - mock_embedding_response - ) + mock_config_instance.transform_embedding_response.return_value = mock_embedding_response # Call the handler result = VertexPassthroughLoggingHandler.vertex_passthrough_handler( @@ -780,26 +742,12 @@ class TestVertexAIPassThroughHandler: ) # Test case 1: Response with textEmbedding should be detected as multimodal - response_with_text_embedding = { - "predictions": [{"textEmbedding": [0.1, 0.2, 0.3]}] - } - assert ( - VertexPassthroughLoggingHandler._is_multimodal_embedding_response( - response_with_text_embedding - ) - is True - ) + response_with_text_embedding = {"predictions": [{"textEmbedding": [0.1, 0.2, 0.3]}]} + assert VertexPassthroughLoggingHandler._is_multimodal_embedding_response(response_with_text_embedding) is True # Test case 2: Response with imageEmbedding should be detected as multimodal - response_with_image_embedding = { - "predictions": [{"imageEmbedding": [0.4, 0.5, 0.6]}] - } - assert ( - VertexPassthroughLoggingHandler._is_multimodal_embedding_response( - response_with_image_embedding - ) - is True - ) + response_with_image_embedding = {"predictions": [{"imageEmbedding": [0.4, 0.5, 0.6]}]} + assert VertexPassthroughLoggingHandler._is_multimodal_embedding_response(response_with_image_embedding) is True # Test case 3: Response with videoEmbeddings should be detected as multimodal response_with_video_embeddings = { @@ -815,43 +763,19 @@ class TestVertexAIPassThroughHandler: } ] } - assert ( - VertexPassthroughLoggingHandler._is_multimodal_embedding_response( - response_with_video_embeddings - ) - is True - ) + assert VertexPassthroughLoggingHandler._is_multimodal_embedding_response(response_with_video_embeddings) is True # Test case 4: Regular text embedding response should NOT be detected as multimodal - regular_embedding_response = { - "predictions": [{"embeddings": {"values": [0.1, 0.2, 0.3]}}] - } - assert ( - VertexPassthroughLoggingHandler._is_multimodal_embedding_response( - regular_embedding_response - ) - is False - ) + regular_embedding_response = {"predictions": [{"embeddings": {"values": [0.1, 0.2, 0.3]}}]} + assert VertexPassthroughLoggingHandler._is_multimodal_embedding_response(regular_embedding_response) is False # Test case 5: Non-embedding response should NOT be detected as multimodal - non_embedding_response = { - "candidates": [{"content": {"parts": [{"text": "Hello world"}]}}] - } - assert ( - VertexPassthroughLoggingHandler._is_multimodal_embedding_response( - non_embedding_response - ) - is False - ) + non_embedding_response = {"candidates": [{"content": {"parts": [{"text": "Hello world"}]}}]} + assert VertexPassthroughLoggingHandler._is_multimodal_embedding_response(non_embedding_response) is False # Test case 6: Empty response should NOT be detected as multimodal empty_response = {} - assert ( - VertexPassthroughLoggingHandler._is_multimodal_embedding_response( - empty_response - ) - is False - ) + assert VertexPassthroughLoggingHandler._is_multimodal_embedding_response(empty_response) is False def test_vertex_passthrough_handler_predict_cost_tracking(self): """ @@ -891,7 +815,9 @@ class TestVertexAIPassThroughHandler: mock_logging_obj.model_call_details = {} # Test URL with /predict endpoint - url_route = "/v1/projects/test-project/locations/us-central1/publishers/google/models/textembedding-gecko@001:predict" + url_route = ( + "/v1/projects/test-project/locations/us-central1/publishers/google/models/textembedding-gecko@001:predict" + ) start_time = datetime.datetime.now() end_time = datetime.datetime.now() @@ -961,7 +887,9 @@ class TestVertexAIPassThroughHandler: mock_logging_obj.litellm_call_id = "test-call-id-embed" mock_logging_obj.model_call_details = {} - url_route = "/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-embedding-001:embedContent" + url_route = ( + "/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-embedding-001:embedContent" + ) start_time = datetime.datetime.now() end_time = datetime.datetime.now() @@ -980,9 +908,7 @@ class TestVertexAIPassThroughHandler: ) assert result is not None - assert ( - result["result"] is not None - ), "result must not be None — logging callbacks need a non-null response" + assert result["result"] is not None, "result must not be None — logging callbacks need a non-null response" assert "kwargs" in result assert result["kwargs"].get("response_cost") == 0.0002 assert result["kwargs"].get("model") == "gemini-embedding-001" @@ -1039,9 +965,7 @@ class TestVertexAIPassThroughHandler: ) assert result is not None - assert ( - result["result"] is not None - ), "result must not be None for batchEmbedContents" + assert result["result"] is not None, "result must not be None for batchEmbedContents" assert result["kwargs"].get("response_cost") == 0.0003 assert result["kwargs"].get("model") == "gemini-embedding-001" assert result["kwargs"].get("custom_llm_provider") == "vertex_ai" @@ -1098,9 +1022,9 @@ class TestVertexAIPassThroughHandler: assert result is not None assert result["result"] is not None - assert ( - result["kwargs"].get("custom_llm_provider") == "gemini" - ), "Google AI Studio embedContent URLs must set custom_llm_provider=gemini, not vertex_ai" + assert result["kwargs"].get("custom_llm_provider") == "gemini", ( + "Google AI Studio embedContent URLs must set custom_llm_provider=gemini, not vertex_ai" + ) assert result["kwargs"].get("model") == "gemini-embedding-2-preview" mock_completion_cost.assert_called_once() @@ -1247,13 +1171,13 @@ class TestVertexAIDiscoveryPassThroughHandler: pass_through_router, ) - endpoint = f"v1/projects/{vertex_project}/locations/{vertex_location}/dataStores/default/servingConfigs/default:search" + endpoint = ( + f"v1/projects/{vertex_project}/locations/{vertex_location}/dataStores/default/servingConfigs/default:search" + ) # Mock request mock_request = Mock() - mock_request.state = ( - None # Prevent Mock from returning a truthy _cached_headers - ) + mock_request.state = None # Prevent Mock from returning a truthy _cached_headers mock_request.method = "POST" mock_request.headers = { "Authorization": "Bearer test-key", @@ -1271,9 +1195,7 @@ class TestVertexAIDiscoveryPassThroughHandler: test_token = "test-auth-token" with ( - mock.patch( - "litellm.llms.vertex_ai.vertex_llm_base.VertexBase.load_auth" - ) as mock_load_auth, + mock.patch("litellm.llms.vertex_ai.vertex_llm_base.VertexBase.load_auth") as mock_load_auth, mock.patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, @@ -1298,9 +1220,7 @@ class TestVertexAIDiscoveryPassThroughHandler: # Mock the discovery handler mock_handler = Mock() - mock_handler.get_default_base_target_url.return_value = ( - "https://discoveryengine.googleapis.com" - ) + mock_handler.get_default_base_target_url.return_value = "https://discoveryengine.googleapis.com" mock_get_handler.return_value = mock_handler # Mock create_pass_through_route to return a function that returns a mock response @@ -1321,10 +1241,7 @@ class TestVertexAIDiscoveryPassThroughHandler: assert test_project in call_args[1]["target"] assert test_location in call_args[1]["target"] assert "Authorization" in call_args[1]["custom_headers"] - assert ( - call_args[1]["custom_headers"]["Authorization"] - == f"Bearer {test_token}" - ) + assert call_args[1]["custom_headers"]["Authorization"] == f"Bearer {test_token}" @pytest.mark.asyncio async def test_vertex_discovery_proxy_route_api_key_auth(self): @@ -1339,17 +1256,13 @@ class TestVertexAIDiscoveryPassThroughHandler: mock_request.method = "POST" mock_response = Mock() - with patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.user_api_key_auth" - ) as mock_auth: + with patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.user_api_key_auth") as mock_auth: mock_auth.return_value = {"api_key": "test-key-123"} with patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_pass_through: - mock_pass_through.return_value = AsyncMock( - return_value={"status": "success"} - ) + mock_pass_through.return_value = AsyncMock(return_value={"status": "success"}) with pytest.raises(HTTPException) as exc_info: await vertex_discovery_proxy_route( @@ -1445,9 +1358,7 @@ async def test_mistral_passthrough_accepts_multipart_without_json_parsing(): assert response == {"ok": True} assert captured_kwargs["is_streaming_request"] is False - assert captured_kwargs["custom_headers"] == { - "Authorization": "Bearer mistral-test-key" - } + assert captured_kwargs["custom_headers"] == {"Authorization": "Bearer mistral-test-key"} class TestBedrockLLMProxyRoute: @@ -1459,9 +1370,7 @@ class TestBedrockLLMProxyRoute: mock_user_api_key_dict = Mock() mock_request_body = {"messages": [{"role": "user", "content": "test"}]} mock_processor = Mock() - mock_processor.base_passthrough_process_llm_request = AsyncMock( - return_value="success" - ) + mock_processor.base_passthrough_process_llm_request = AsyncMock(return_value="success") with ( patch( @@ -1473,9 +1382,10 @@ class TestBedrockLLMProxyRoute: return_value=mock_processor, ), ): - # Test application-inference-profile endpoint - endpoint = "model/arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/r742sbn2zckd/converse" + endpoint = ( + "model/arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/r742sbn2zckd/converse" + ) result = await bedrock_llm_proxy_route( endpoint=endpoint, @@ -1485,9 +1395,7 @@ class TestBedrockLLMProxyRoute: ) mock_processor.base_passthrough_process_llm_request.assert_called_once() - call_kwargs = ( - mock_processor.base_passthrough_process_llm_request.call_args.kwargs - ) + call_kwargs = mock_processor.base_passthrough_process_llm_request.call_args.kwargs # For application-inference-profile, model should be "arn:aws:bedrock:us-east-1:026090525607:application-inference-profile/r742sbn2zckd" assert ( @@ -1504,9 +1412,7 @@ class TestBedrockLLMProxyRoute: mock_user_api_key_dict = Mock() mock_request_body = {"messages": [{"role": "user", "content": "test"}]} mock_processor = Mock() - mock_processor.base_passthrough_process_llm_request = AsyncMock( - return_value="success" - ) + mock_processor.base_passthrough_process_llm_request = AsyncMock(return_value="success") with ( patch( @@ -1518,7 +1424,6 @@ class TestBedrockLLMProxyRoute: return_value=mock_processor, ), ): - # Test regular model endpoint endpoint = "model/anthropic.claude-3-sonnet-20240229-v1:0/converse" @@ -1529,9 +1434,7 @@ class TestBedrockLLMProxyRoute: user_api_key_dict=mock_user_api_key_dict, ) mock_processor.base_passthrough_process_llm_request.assert_called_once() - call_kwargs = ( - mock_processor.base_passthrough_process_llm_request.call_args.kwargs - ) + call_kwargs = mock_processor.base_passthrough_process_llm_request.call_args.kwargs # For regular models, model should be just the model ID assert call_kwargs["model"] == "anthropic.claude-3-sonnet-20240229-v1:0" @@ -1554,9 +1457,7 @@ class TestBedrockLLMProxyRoute: # Create a mock httpx.Response for the error mock_error_response = Mock(spec=httpx.Response) mock_error_response.status_code = 400 - mock_error_response.aread = AsyncMock( - return_value=bedrock_error_message.encode("utf-8") - ) + mock_error_response.aread = AsyncMock(return_value=bedrock_error_message.encode("utf-8")) # Create the HTTPStatusError mock_http_error = httpx.HTTPStatusError( @@ -1573,9 +1474,7 @@ class TestBedrockLLMProxyRoute: mock_request.url = MagicMock() mock_request.url.path = "/bedrock/model/test-model/converse" - mock_request_body = { - "messages": [{"role": "user", "content": [{"textaaa": "Hello"}]}] - } + mock_request_body = {"messages": [{"role": "user", "content": [{"textaaa": "Hello"}]}]} mock_llm_router = Mock() @@ -1616,9 +1515,8 @@ class TestBedrockLLMProxyRoute: ) assert exc_info.value.status_code == 400 - assert ( - "ContentBlock object at messages.0.content.0 must set one of the following keys" - in str(exc_info.value.detail) + assert "ContentBlock object at messages.0.content.0 must set one of the following keys" in str( + exc_info.value.detail ) @pytest.mark.asyncio @@ -1696,24 +1594,14 @@ class TestBedrockLLMProxyRoute: deployment_litellm_params = deployment.get("litellm_params", {}) # Verify model-specific credentials are in the deployment - assert ( - deployment_litellm_params.get("aws_access_key_id") == model_access_key - ) - assert ( - deployment_litellm_params.get("aws_secret_access_key") - == model_secret_key - ) + assert deployment_litellm_params.get("aws_access_key_id") == model_access_key + assert deployment_litellm_params.get("aws_secret_access_key") == model_secret_key assert deployment_litellm_params.get("aws_region_name") == model_region - assert ( - deployment_litellm_params.get("aws_session_token") - == model_session_token - ) + assert deployment_litellm_params.get("aws_session_token") == model_session_token # Verify environment variables are NOT in the deployment assert deployment_litellm_params.get("aws_access_key_id") != env_access_key - assert ( - deployment_litellm_params.get("aws_secret_access_key") != env_secret_key - ) + assert deployment_litellm_params.get("aws_secret_access_key") != env_secret_key assert deployment_litellm_params.get("aws_region_name") != env_region # Test 3: Verify credentials are passed through the passthrough route @@ -1724,9 +1612,7 @@ class TestBedrockLLMProxyRoute: captured_kwargs.update(kwargs) mock_response = MagicMock() mock_response.status_code = 200 - mock_response.aread = AsyncMock( - return_value=b'{"content": [{"text": "Hello"}]}' - ) + mock_response.aread = AsyncMock(return_value=b'{"content": [{"text": "Hello"}]}') return mock_response mock_request = MagicMock(spec=Request) @@ -1736,9 +1622,7 @@ class TestBedrockLLMProxyRoute: mock_request.url = MagicMock() mock_request.url.path = "/bedrock/model/claude-opus-4-1/converse" - mock_request_body = { - "messages": [{"role": "user", "content": [{"text": "Hello"}]}] - } + mock_request_body = {"messages": [{"role": "user", "content": [{"text": "Hello"}]}]} mock_user_api_key_dict = Mock() mock_user_api_key_dict.api_key = "test-key" @@ -1759,9 +1643,7 @@ class TestBedrockLLMProxyRoute: # Setup mock response mock_response = MagicMock() mock_response.status_code = 200 - mock_response.aread = AsyncMock( - return_value=b'{"content": [{"text": "Hello"}]}' - ) + mock_response.aread = AsyncMock(return_value=b'{"content": [{"text": "Hello"}]}') mock_process.return_value = mock_response # Call the handler @@ -1968,9 +1850,7 @@ class TestLLMPassthroughFactoryProxyRoute: mock_user_api_key_dict = MagicMock() with ( - patch( - "litellm.utils.ProviderConfigManager.get_provider_model_info" - ) as mock_get_provider, + patch("litellm.utils.ProviderConfigManager.get_provider_model_info") as mock_get_provider, patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials" ) as mock_get_creds, @@ -1980,9 +1860,7 @@ class TestLLMPassthroughFactoryProxyRoute: ): mock_provider_config = MagicMock() mock_provider_config.get_api_base.return_value = "https://example.com/v1" - mock_provider_config.validate_environment.return_value = { - "x-api-key": "dummy" - } + mock_provider_config.validate_environment.return_value = {"x-api-key": "dummy"} mock_get_provider.return_value = mock_provider_config mock_get_creds.return_value = "dummy" @@ -1998,12 +1876,8 @@ class TestLLMPassthroughFactoryProxyRoute: ) assert result == "success" - mock_get_provider.assert_called_once_with( - provider=litellm.LlmProviders(LlmProviders.VLLM), model=None - ) - mock_get_creds.assert_called_once_with( - custom_llm_provider=LlmProviders.VLLM, region_name=None - ) + mock_get_provider.assert_called_once_with(provider=litellm.LlmProviders(LlmProviders.VLLM), model=None) + mock_get_creds.assert_called_once_with(custom_llm_provider=LlmProviders.VLLM, region_name=None) mock_create_route.assert_called_once_with( endpoint="/chat/completions", target="https://example.com/v1/chat/completions", @@ -2023,10 +1897,10 @@ class TestVLLMProxyRoute: "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model", return_value=True, ) - @patch("litellm.proxy.proxy_server.llm_router") # test-quality-ok: patching litellm internal for unit test isolation - async def test_vllm_proxy_route_with_router_model( - self, mock_llm_router, mock_is_router, mock_get_body - ): + @patch( + "litellm.proxy.proxy_server.llm_router" + ) # test-quality-ok: patching litellm internal for unit test isolation + async def test_vllm_proxy_route_with_router_model(self, mock_llm_router, mock_is_router, mock_get_body): mock_request = MagicMock(spec=Request) mock_request.method = "POST" mock_request.headers = {"content-type": "application/json"} @@ -2059,9 +1933,7 @@ class TestVLLMProxyRoute: @patch( # test-quality-ok: patching litellm internal for unit test isolation "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.llm_passthrough_factory_proxy_route" ) - async def test_vllm_proxy_route_fallback_to_factory( - self, mock_factory_route, mock_is_router, mock_get_body - ): + async def test_vllm_proxy_route_fallback_to_factory(self, mock_factory_route, mock_is_router, mock_get_body): mock_request = MagicMock(spec=Request) mock_fastapi_response = MagicMock(spec=Response) mock_user_api_key_dict = MagicMock() @@ -2088,10 +1960,10 @@ class TestGigachatProxyRoute: "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model", return_value=True, ) - @patch("litellm.proxy.proxy_server.llm_router") # test-quality-ok: patching litellm internal for unit test isolation - async def test_gigachat_proxy_route_with_router_model( - self, mock_llm_router, mock_is_router, mock_get_body - ): + @patch( + "litellm.proxy.proxy_server.llm_router" + ) # test-quality-ok: patching litellm internal for unit test isolation + async def test_gigachat_proxy_route_with_router_model(self, mock_llm_router, mock_is_router, mock_get_body): mock_request = MagicMock(spec=Request) mock_request.method = "POST" mock_request.headers = {"content-type": "application/json"} @@ -2344,21 +2216,25 @@ class TestGigachatProxyRoute: return _inner() - with patch.object( - processor, - "common_processing_pre_call_logic", - new=AsyncMock( - return_value=( - processor.data, - processor.data["litellm_logging_obj"], - ) + with ( + patch.object( + processor, + "common_processing_pre_call_logic", + new=AsyncMock( + return_value=( + processor.data, + processor.data["litellm_logging_obj"], + ) + ), + ), + patch( # test-quality-ok: patching litellm internal for unit test isolation + "litellm.proxy.common_request_processing.route_request", + new=_fake_route_request, + ), + patch( # test-quality-ok: patching litellm internal for unit test isolation + "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing.get_custom_headers", + return_value={"x-litellm-call-id": "call-123"}, ), - ), patch( # test-quality-ok: patching litellm internal for unit test isolation - "litellm.proxy.common_request_processing.route_request", - new=_fake_route_request, - ), patch( # test-quality-ok: patching litellm internal for unit test isolation - "litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing.get_custom_headers", - return_value={"x-litellm-call-id": "call-123"}, ): result = await processor.base_passthrough_process_llm_request( request=mock_request, @@ -2401,9 +2277,7 @@ class TestForwardHeaders: # Create a mock request with custom headers mock_request = MagicMock(spec=Request) - mock_request.state = ( - None # Prevent MagicMock from returning a truthy _cached_headers - ) + mock_request.state = None # Prevent MagicMock from returning a truthy _cached_headers mock_request.method = "POST" mock_request.url = MagicMock() mock_request.url.path = "/test/endpoint" @@ -2438,9 +2312,7 @@ class TestForwardHeaders: mock_httpx_response = MagicMock() mock_httpx_response.status_code = 200 mock_httpx_response.headers = {"content-type": "application/json"} - mock_httpx_response.aiter_bytes = AsyncMock( - return_value=[b'{"result": "success"}'] - ) + mock_httpx_response.aiter_bytes = AsyncMock(return_value=[b'{"result": "success"}']) mock_httpx_response.aread = AsyncMock(return_value=b'{"result": "success"}') with ( @@ -2465,9 +2337,7 @@ class TestForwardHeaders: mock_logging_obj.pre_call_hook = AsyncMock(return_value=mock_request_body) mock_logging_obj.post_call_success_hook = AsyncMock() mock_logging_obj.post_call_failure_hook = AsyncMock() - mock_logging_obj.post_call_response_headers_hook = AsyncMock( - return_value={} - ) + mock_logging_obj.post_call_response_headers_hook = AsyncMock(return_value={}) # Call pass_through_request with forward_headers=True result = await pass_through_request( @@ -2540,9 +2410,7 @@ class TestForwardHeaders: mock_httpx_response = MagicMock() mock_httpx_response.status_code = 200 mock_httpx_response.headers = {"content-type": "application/json"} - mock_httpx_response.aiter_bytes = AsyncMock( - return_value=[b'{"result": "success"}'] - ) + mock_httpx_response.aiter_bytes = AsyncMock(return_value=[b'{"result": "success"}']) mock_httpx_response.aread = AsyncMock(return_value=b'{"result": "success"}') with ( @@ -2567,9 +2435,7 @@ class TestForwardHeaders: mock_logging_obj.pre_call_hook = AsyncMock(return_value=mock_request_body) mock_logging_obj.post_call_success_hook = AsyncMock() mock_logging_obj.post_call_failure_hook = AsyncMock() - mock_logging_obj.post_call_response_headers_hook = AsyncMock( - return_value={} - ) + mock_logging_obj.post_call_response_headers_hook = AsyncMock(return_value={}) # Call pass_through_request with forward_headers=False (default) result = await pass_through_request( @@ -2627,15 +2493,11 @@ class TestForwardHeaders: mock_httpx_response = MagicMock() mock_httpx_response.status_code = 200 mock_httpx_response.headers = {"content-type": "application/json"} - mock_httpx_response.aiter_bytes = AsyncMock( - return_value=[b'{"result": "success"}'] - ) + mock_httpx_response.aiter_bytes = AsyncMock(return_value=[b'{"result": "success"}']) mock_httpx_response.aread = AsyncMock(return_value=b'{"result": "success"}') with ( - patch( - "litellm.utils.ProviderConfigManager.get_provider_model_info" - ) as mock_get_provider, + patch("litellm.utils.ProviderConfigManager.get_provider_model_info") as mock_get_provider, patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials" ) as mock_get_creds, @@ -2651,9 +2513,7 @@ class TestForwardHeaders: # Setup provider config mock_provider_config = MagicMock() mock_provider_config.get_api_base.return_value = "https://api.openai.com/v1" - mock_provider_config.validate_environment.return_value = { - "authorization": "Bearer sk-test" - } + mock_provider_config.validate_environment.return_value = {"authorization": "Bearer sk-test"} mock_get_provider.return_value = mock_provider_config mock_get_creds.return_value = "sk-test" @@ -2665,9 +2525,7 @@ class TestForwardHeaders: mock_get_client.return_value = mock_client_obj # Setup mock logging object - mock_logging_obj.pre_call_hook = AsyncMock( - return_value={"messages": [{"role": "user", "content": "test"}]} - ) + mock_logging_obj.pre_call_hook = AsyncMock(return_value={"messages": [{"role": "user", "content": "test"}]}) mock_logging_obj.post_call_success_hook = AsyncMock() # This is the key part - when create_pass_through_route is called with _forward_headers=True @@ -2758,24 +2616,16 @@ class TestMilvusProxyRoute: ): # Setup mocks mock_provider_config = MagicMock() - mock_provider_config.get_auth_credentials.return_value = { - "headers": {"Authorization": "Bearer test-token"} - } + mock_provider_config.get_auth_credentials.return_value = {"headers": {"Authorization": "Bearer test-token"}} mock_provider_config.get_complete_url.return_value = api_base mock_get_config.return_value = mock_provider_config mock_index_registry.is_vector_store_index.return_value = True - mock_index_registry.get_vector_store_index_by_name.return_value = ( - mock_index_object - ) + mock_index_registry.get_vector_store_index_by_name.return_value = mock_index_object - mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = ( - mock_vector_store - ) + mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = mock_vector_store - mock_endpoint_func = AsyncMock( - return_value={"results": [{"id": 1, "distance": 0.5}]} - ) + mock_endpoint_func = AsyncMock(return_value={"results": [{"id": 1, "distance": 0.5}]}) mock_create_route.return_value = mock_endpoint_func # Call the route @@ -2788,9 +2638,7 @@ class TestMilvusProxyRoute: # Verify calls mock_get_body.assert_called_once() - mock_index_registry.is_vector_store_index.assert_called_once_with( - vector_store_index_name=collection_name - ) + mock_index_registry.is_vector_store_index.assert_called_once_with(vector_store_index_name=collection_name) mock_is_allowed.assert_called_once() mock_safe_set.assert_called_once() @@ -2802,9 +2650,7 @@ class TestMilvusProxyRoute: mock_create_route.assert_called_once() create_route_args = mock_create_route.call_args[1] assert "vectors/search" in create_route_args["target"] - assert create_route_args["custom_headers"] == { - "Authorization": "Bearer test-token" - } + assert create_route_args["custom_headers"] == {"Authorization": "Bearer test-token"} # Verify endpoint function was called mock_endpoint_func.assert_awaited_once() @@ -2817,7 +2663,6 @@ class TestMilvusProxyRoute: """ from fastapi import HTTPException - mock_request = MagicMock(spec=Request) mock_response = MagicMock(spec=Response) mock_user_api_key_dict = MagicMock() @@ -2851,7 +2696,6 @@ class TestMilvusProxyRoute: """ from fastapi import HTTPException - mock_request = MagicMock(spec=Request) mock_response = MagicMock(spec=Response) mock_user_api_key_dict = MagicMock() @@ -2869,9 +2713,7 @@ class TestMilvusProxyRoute: ) assert exc_info.value.status_code == 500 - assert "Unable to find Milvus vector store config" in str( - exc_info.value.detail - ) + assert "Unable to find Milvus vector store config" in str(exc_info.value.detail) @pytest.mark.asyncio async def test_milvus_proxy_route_no_index_registry(self): @@ -2880,7 +2722,6 @@ class TestMilvusProxyRoute: """ from fastapi import HTTPException - collection_name = "test-collection" mock_request = MagicMock(spec=Request) @@ -2908,9 +2749,7 @@ class TestMilvusProxyRoute: ) assert exc_info.value.status_code == 500 - assert "Unable to find Milvus vector store index registry" in str( - exc_info.value.detail - ) + assert "Unable to find Milvus vector store index registry" in str(exc_info.value.detail) @pytest.mark.asyncio async def test_milvus_proxy_route_not_managed_index(self): @@ -2919,7 +2758,6 @@ class TestMilvusProxyRoute: """ from fastapi import HTTPException - collection_name = "unmanaged-collection" mock_request = MagicMock(spec=Request) @@ -2949,9 +2787,8 @@ class TestMilvusProxyRoute: ) assert exc_info.value.status_code == 400 - assert ( - f"Collection {collection_name} is not a litellm managed vector store index" - in str(exc_info.value.detail) + assert f"Collection {collection_name} is not a litellm managed vector store index" in str( + exc_info.value.detail ) @pytest.mark.asyncio @@ -2983,22 +2820,16 @@ class TestMilvusProxyRoute: patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint" ), - patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body" - ), + patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body"), patch.object(litellm, "vector_store_index_registry") as mock_index_registry, patch.object(litellm, "vector_store_registry") as mock_vector_registry, ): mock_get_config.return_value = MagicMock() mock_index_registry.is_vector_store_index.return_value = True - mock_index_registry.get_vector_store_index_by_name.return_value = ( - mock_index_object - ) - mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = ( - None - ) + mock_index_registry.get_vector_store_index_by_name.return_value = mock_index_object + mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = None - with pytest.raises(Exception, match='Vector store not found for missing-store') as exc_info: + with pytest.raises(Exception, match="Vector store not found for missing-store") as exc_info: await milvus_proxy_route( endpoint="vectors/search", request=mock_request, @@ -3006,9 +2837,7 @@ class TestMilvusProxyRoute: user_api_key_dict=mock_user_api_key_dict, ) - assert f"Vector store not found for {vector_store_name}" in str( - exc_info.value - ) + assert f"Vector store not found for {vector_store_name}" in str(exc_info.value) @pytest.mark.asyncio async def test_milvus_proxy_route_no_api_base(self): @@ -3041,9 +2870,7 @@ class TestMilvusProxyRoute: patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint" ), - patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body" - ), + patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body"), patch.object(litellm, "vector_store_index_registry") as mock_index_registry, patch.object(litellm, "vector_store_registry") as mock_vector_registry, ): @@ -3053,14 +2880,10 @@ class TestMilvusProxyRoute: mock_get_config.return_value = mock_provider_config mock_index_registry.is_vector_store_index.return_value = True - mock_index_registry.get_vector_store_index_by_name.return_value = ( - mock_index_object - ) - mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = ( - mock_vector_store - ) + mock_index_registry.get_vector_store_index_by_name.return_value = mock_index_object + mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = mock_vector_store - with pytest.raises(Exception, match='api_base not found in vector store configuration for') as exc_info: + with pytest.raises(Exception, match="api_base not found in vector store configuration for") as exc_info: await milvus_proxy_route( endpoint="vectors/search", request=mock_request, @@ -3068,10 +2891,7 @@ class TestMilvusProxyRoute: user_api_key_dict=mock_user_api_key_dict, ) - assert ( - f"api_base not found in vector store configuration for {vector_store_name}" - in str(exc_info.value) - ) + assert f"api_base not found in vector store configuration for {vector_store_name}" in str(exc_info.value) @pytest.mark.asyncio async def test_milvus_proxy_route_endpoint_without_leading_slash(self): @@ -3105,9 +2925,7 @@ class TestMilvusProxyRoute: patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint" ), - patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body" - ), + patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body"), patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, @@ -3120,12 +2938,8 @@ class TestMilvusProxyRoute: mock_get_config.return_value = mock_provider_config mock_index_registry.is_vector_store_index.return_value = True - mock_index_registry.get_vector_store_index_by_name.return_value = ( - mock_index_object - ) - mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = ( - mock_vector_store - ) + mock_index_registry.get_vector_store_index_by_name.return_value = mock_index_object + mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = mock_vector_store mock_endpoint_func = AsyncMock(return_value={"status": "success"}) mock_create_route.return_value = mock_endpoint_func @@ -3173,9 +2987,7 @@ class TestOpenAIPassthroughRoute: "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, ): - mock_endpoint_func = AsyncMock( - return_value={"id": "resp_123", "status": "completed"} - ) + mock_endpoint_func = AsyncMock(return_value={"id": "resp_123", "status": "completed"}) mock_create_route.return_value = mock_endpoint_func # Call the route with /v1/responses endpoint @@ -3223,9 +3035,7 @@ class TestOpenAIPassthroughRoute: "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, ): - mock_endpoint_func = AsyncMock( - return_value={"id": "chatcmpl-123", "choices": []} - ) + mock_endpoint_func = AsyncMock(return_value={"id": "chatcmpl-123", "choices": []}) mock_create_route.return_value = mock_endpoint_func result = await openai_proxy_route( @@ -3291,9 +3101,7 @@ class TestOpenAIPassthroughRoute: "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, ): - mock_endpoint_func = AsyncMock( - return_value={"id": "asst_123", "object": "assistant"} - ) + mock_endpoint_func = AsyncMock(return_value={"id": "asst_123", "object": "assistant"}) mock_create_route.return_value = mock_endpoint_func result = await openai_proxy_route( @@ -3398,9 +3206,7 @@ class TestCursorProxyRoute: "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, ): - mock_endpoint_func = AsyncMock( - return_value={"agents": [], "nextCursor": None} - ) + mock_endpoint_func = AsyncMock(return_value={"agents": [], "nextCursor": None}) mock_create_route.return_value = mock_endpoint_func result = await cursor_proxy_route( @@ -3414,12 +3220,8 @@ class TestCursorProxyRoute: call_args = mock_create_route.call_args[1] assert call_args["target"] == "https://api.cursor.com/v0/agents" - expected_auth = base64.b64encode(f"{test_api_key}:".encode("utf-8")).decode( - "ascii" - ) - assert ( - call_args["custom_headers"]["Authorization"] == f"Basic {expected_auth}" - ) + expected_auth = base64.b64encode(f"{test_api_key}:".encode("utf-8")).decode("ascii") + assert call_args["custom_headers"]["Authorization"] == f"Basic {expected_auth}" assert result == {"agents": [], "nextCursor": None} @@ -3443,7 +3245,7 @@ class TestCursorProxyRoute: [], ), ): - with pytest.raises(Exception, match='Cursor API key not found\\. Add Cursor credentials via') as exc_info: + with pytest.raises(Exception, match="Cursor API key not found\\. Add Cursor credentials via") as exc_info: await cursor_proxy_route( endpoint="v0/agents", request=mock_request, @@ -3502,9 +3304,7 @@ class TestCursorProxyRoute: import base64 expected_auth = base64.b64encode(b"crsr_ui_test_key:").decode("ascii") - assert ( - call_args["custom_headers"]["Authorization"] == f"Basic {expected_auth}" - ) + assert call_args["custom_headers"]["Authorization"] == f"Basic {expected_auth}" @pytest.mark.asyncio async def test_cursor_proxy_route_custom_api_base(self): @@ -3517,9 +3317,7 @@ class TestCursorProxyRoute: mock_user_api_key_dict = MagicMock() with ( - patch.dict( - os.environ, {"CURSOR_API_BASE": "https://custom-cursor.example.com"} - ), + patch.dict(os.environ, {"CURSOR_API_BASE": "https://custom-cursor.example.com"}), patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials", return_value="test-key", @@ -3599,12 +3397,10 @@ class TestVertexRawPredictStreamingClassification: """ RAW_PREDICT_ENDPOINT = ( - "v1/projects/test-project/locations/us-east5/publishers/anthropic/models/" - "claude-sonnet-4-6:streamRawPredict" + "v1/projects/test-project/locations/us-east5/publishers/anthropic/models/claude-sonnet-4-6:streamRawPredict" ) GENERATE_CONTENT_ENDPOINT = ( - "v1/projects/test-project/locations/us-east5/publishers/google/models/" - "gemini-2.5-flash:streamGenerateContent" + "v1/projects/test-project/locations/us-east5/publishers/google/models/gemini-2.5-flash:streamGenerateContent" ) async def _capture_passthrough_kwargs(self, endpoint: str, body: object) -> dict: @@ -3789,10 +3585,7 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: """ VKEY = "sk-litellm-victim-key" - ENDPOINT = ( - "v1/projects/my-proj/locations/us-central1/publishers/google/models/" - "gemini-2.5-flash:generateContent" - ) + ENDPOINT = "v1/projects/my-proj/locations/us-central1/publishers/google/models/gemini-2.5-flash:generateContent" async def _run( self, @@ -3920,7 +3713,9 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: (b"content-type", b"application/json"), ], ) - assert forwarded is None, f"a virtual key echoed as '{scheme} ' in Authorization must be stripped, not forwarded" + assert forwarded is None, ( + f"a virtual key echoed as '{scheme} ' in Authorization must be stripped, not forwarded" + ) assert raised is not None and raised.status_code == 401 @pytest.mark.asyncio @@ -3966,8 +3761,7 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: @pytest.mark.parametrize( "credential_header", sorted( - SpecialHeaders.litellm_credential_header_names() - - {"authorization", "x-goog-api-key", "x-litellm-api-key"} + SpecialHeaders.litellm_credential_header_names() - {"authorization", "x-goog-api-key", "x-litellm-api-key"} ), ) async def test_every_non_google_credential_header_is_dropped_by_name(self, monkeypatch, credential_header): @@ -4042,7 +3836,9 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: assert raised is not None and raised.status_code == 401 @pytest.mark.asyncio - async def test_authenticated_authorization_is_stripped_over_a_lower_precedence_pass_through_header(self, monkeypatch): + async def test_authenticated_authorization_is_stripped_over_a_lower_precedence_pass_through_header( + self, monkeypatch + ): with mock.patch.dict( # test-quality-ok: general_settings is the real proxy config surface for pass_through_endpoints; no injection seam exists on this route "litellm.proxy.proxy_server.general_settings", {"pass_through_endpoints": [{"headers": {"litellm_user_api_key": "x-company-key"}}]}, @@ -4059,7 +3855,9 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: assert raised is None assert forwarded is not None assert forwarded.get("x-goog-api-key") == "AIza-real-google-api-key" - assert "authorization" not in forwarded, "Authorization authenticated (higher precedence) so its key must be stripped" + assert "authorization" not in forwarded, ( + "Authorization authenticated (higher precedence) so its key must be stripped" + ) assert "x-company-key" not in forwarded assert self.VKEY not in " ".join(f"{name}:{value}" for name, value in forwarded.items()) @@ -4090,7 +3888,9 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: (b"content-type", b"application/json"), ], ) - assert forwarded is None, "a virtual key in the mapped-route litellm_user_api_key header must be dropped, not forwarded" + assert forwarded is None, ( + "a virtual key in the mapped-route litellm_user_api_key header must be dropped, not forwarded" + ) assert raised is not None and raised.status_code == 401 GOOGLE_OAUTH_TOKEN = "ya29.byo-google-oauth-token" @@ -4142,7 +3942,9 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: @pytest.mark.parametrize( ("credential", "authenticated"), [ - pytest.param("modified_key", UserAPIKeyAuth(api_key="modified_key"), id="custom-auth-echoing-opaque-credential"), + pytest.param( + "modified_key", UserAPIKeyAuth(api_key="modified_key"), id="custom-auth-echoing-opaque-credential" + ), pytest.param( LITELLM_JWT, UserAPIKeyAuth(api_key=LITELLM_JWT, user_id="jwt-subject"), @@ -4200,7 +4002,9 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: raised, forwarded = await self._run( monkeypatch, [(b"authorization", b"Bearer sk-master-1234"), (b"content-type", b"application/json")], - authenticated=UserAPIKeyAuth(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN), + authenticated=UserAPIKeyAuth( + api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN + ), ) assert forwarded is None, "the master key must never reach the upstream forwarder" assert raised is not None and raised.status_code == 401 @@ -4214,7 +4018,9 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: (b"x-goog-api-key", b"AIza-real-google-api-key"), (b"content-type", b"application/json"), ], - authenticated=UserAPIKeyAuth(api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN), + authenticated=UserAPIKeyAuth( + api_key=LITELLM_PROXY_MASTER_KEY_ALIAS, user_role=LitellmUserRoles.PROXY_ADMIN + ), ) assert raised is None assert forwarded is not None @@ -4787,9 +4593,7 @@ class TestVertexAILiveWebsocketPassthrough: ] ) monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router) - monkeypatch.setattr( - passthrough_module.passthrough_endpoint_router, "default_vertex_config", None - ) + monkeypatch.setattr(passthrough_module.passthrough_endpoint_router, "default_vertex_config", None) self._clear_vertex_env(monkeypatch) websocket = self._websocket() ensure_token = AsyncMock(return_value=("token-abc", "proj-db")) @@ -4931,9 +4735,7 @@ class TestVertexAILiveWebsocketPassthrough: ) monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) - monkeypatch.setattr( - passthrough_module.passthrough_endpoint_router, "default_vertex_config", None - ) + monkeypatch.setattr(passthrough_module.passthrough_endpoint_router, "default_vertex_config", None) self._clear_vertex_env(monkeypatch) websocket = self._websocket() ensure_token = AsyncMock(side_effect=Exception("Unable to find your credentials")) @@ -5219,6 +5021,42 @@ class TestTypeSafePassthroughRoute: request.json = AsyncMock(return_value=body) return request + @pytest.fixture + def client(self, monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]: + from litellm.proxy.proxy_server import app + + monkeypatch.setenv("TYPESAFE_API_KEY", "typesafe-test-key") + monkeypatch.setenv("TYPESAFE_API_BASE", "https://typesafe.example/base") + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + monkeypatch.setitem(app.dependency_overrides, user_api_key_auth, lambda: UserAPIKeyAuth(api_key="sk-virtual")) + yield TestClient(app) + + @pytest.mark.parametrize( + "method, body", + [ + ("GET", None), + ("POST", {"state": "x"}), + ("PUT", {"state": "x"}), + ("DELETE", None), + ("PATCH", {"state": "x"}), + ], + ) + def test_forwards_every_method_and_body_upstream( + self, client: TestClient, method: str, body: dict[str, str] | None + ) -> None: + with respx.mock(assert_all_called=True) as upstream: + route = upstream.request(method, "https://typesafe.example/base/v1/systemone").mock( + return_value=httpx.Response(200, json={"id": "upstream_123"}) + ) + response = client.request(method, "/typesafe/v1/systemone", json=body) + + assert (response.status_code, response.json()) == (200, {"id": "upstream_123"}) + sent: Final = route.calls.last.request + assert sent.headers["authorization"] == "Bearer typesafe-test-key" + assert json.loads(sent.content or b"{}") == (body or {}) + @pytest.mark.asyncio async def test_forwards_target_auth_headers_provider_and_query(self, monkeypatch): monkeypatch.setenv("TYPESAFE_API_KEY", "typesafe-test-key") diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index cff170ad27e..e20aa0b39dd 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -16180,16 +16180,28 @@ export interface paths { * @description [Docs](https://docs.litellm.ai/docs/pass_through/typesafe) */ get: operations["typesafe_proxy_route_typesafe__endpoint__get"]; - put?: never; + /** + * Typesafe Proxy Route + * @description [Docs](https://docs.litellm.ai/docs/pass_through/typesafe) + */ + put: operations["typesafe_proxy_route_typesafe__endpoint__put"]; /** * Typesafe Proxy Route * @description [Docs](https://docs.litellm.ai/docs/pass_through/typesafe) */ post: operations["typesafe_proxy_route_typesafe__endpoint__post"]; - delete?: never; + /** + * Typesafe Proxy Route + * @description [Docs](https://docs.litellm.ai/docs/pass_through/typesafe) + */ + delete: operations["typesafe_proxy_route_typesafe__endpoint__delete"]; options?: never; head?: never; - patch?: never; + /** + * Typesafe Proxy Route + * @description [Docs](https://docs.litellm.ai/docs/pass_through/typesafe) + */ + patch: operations["typesafe_proxy_route_typesafe__endpoint__patch"]; trace?: never; }; "/update/default_team_settings": { @@ -59701,6 +59713,37 @@ export interface operations { }; }; }; + typesafe_proxy_route_typesafe__endpoint__put: { + parameters: { + query?: never; + header?: never; + path: { + endpoint: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; typesafe_proxy_route_typesafe__endpoint__post: { parameters: { query?: never; @@ -59732,6 +59775,68 @@ export interface operations { }; }; }; + typesafe_proxy_route_typesafe__endpoint__delete: { + parameters: { + query?: never; + header?: never; + path: { + endpoint: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + typesafe_proxy_route_typesafe__endpoint__patch: { + parameters: { + query?: never; + header?: never; + path: { + endpoint: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; update_default_team_settings_update_default_team_settings_patch: { parameters: { query?: never;