fix: multiple logs

This commit is contained in:
Harshit28j 2026-03-06 02:14:51 +05:30
parent 2b91978b99
commit af8431539e
3 changed files with 308 additions and 161 deletions

View file

@ -352,9 +352,9 @@ class Logging(LiteLLMLoggingBaseClass):
)
self.function_id = function_id
self.streaming_chunks: List[Any] = [] # for generating complete stream response
self.sync_streaming_chunks: List[Any] = (
[]
) # for generating complete stream response
self.sync_streaming_chunks: List[
Any
] = [] # for generating complete stream response
self.log_raw_request_response = log_raw_request_response
# Initialize dynamic callbacks
@ -746,9 +746,9 @@ class Logging(LiteLLMLoggingBaseClass):
prompt_spec=prompt_spec,
dynamic_callback_params=dynamic_callback_params,
):
self.model_call_details["prompt_integration"] = (
logger.__class__.__name__
)
self.model_call_details[
"prompt_integration"
] = logger.__class__.__name__
return logger
except Exception:
# If check fails, continue to next logger
@ -816,9 +816,9 @@ class Logging(LiteLLMLoggingBaseClass):
if anthropic_cache_control_logger := AnthropicCacheControlHook.get_custom_logger_for_anthropic_cache_control_hook(
non_default_params
):
self.model_call_details["prompt_integration"] = (
anthropic_cache_control_logger.__class__.__name__
)
self.model_call_details[
"prompt_integration"
] = anthropic_cache_control_logger.__class__.__name__
return anthropic_cache_control_logger
#########################################################
@ -830,9 +830,9 @@ class Logging(LiteLLMLoggingBaseClass):
internal_usage_cache=None,
llm_router=None,
)
self.model_call_details["prompt_integration"] = (
vector_store_custom_logger.__class__.__name__
)
self.model_call_details[
"prompt_integration"
] = vector_store_custom_logger.__class__.__name__
# Add to global callbacks so post-call hooks are invoked
if (
vector_store_custom_logger
@ -892,9 +892,9 @@ class Logging(LiteLLMLoggingBaseClass):
model
): # if model name was changes pre-call, overwrite the initial model call name with the new one
self.model_call_details["model"] = model
self.model_call_details["litellm_params"]["api_base"] = (
self._get_masked_api_base(additional_args.get("api_base", ""))
)
self.model_call_details["litellm_params"][
"api_base"
] = self._get_masked_api_base(additional_args.get("api_base", ""))
def pre_call(self, input, api_key, model=None, additional_args={}): # noqa: PLR0915
# Log the exact input to the LLM API
@ -923,10 +923,10 @@ class Logging(LiteLLMLoggingBaseClass):
try:
# [Non-blocking Extra Debug Information in metadata]
if turn_off_message_logging is True:
_metadata["raw_request"] = (
"redacted by litellm. \
_metadata[
"raw_request"
] = "redacted by litellm. \
'litellm.turn_off_message_logging=True'"
)
else:
curl_command = self._get_request_curl_command(
api_base=additional_args.get("api_base", ""),
@ -937,34 +937,34 @@ class Logging(LiteLLMLoggingBaseClass):
_metadata["raw_request"] = str(curl_command)
# split up, so it's easier to parse in the UI
self.model_call_details["raw_request_typed_dict"] = (
RawRequestTypedDict(
raw_request_api_base=str(
additional_args.get("api_base") or ""
),
raw_request_body=self._get_raw_request_body(
additional_args.get("complete_input_dict", {})
),
# NOTE: setting ignore_sensitive_headers to True will cause
# the Authorization header to be leaked when calls to the health
# endpoint are made and fail.
raw_request_headers=self._get_masked_headers(
additional_args.get("headers", {}) or {},
),
error=None,
)
self.model_call_details[
"raw_request_typed_dict"
] = RawRequestTypedDict(
raw_request_api_base=str(
additional_args.get("api_base") or ""
),
raw_request_body=self._get_raw_request_body(
additional_args.get("complete_input_dict", {})
),
# NOTE: setting ignore_sensitive_headers to True will cause
# the Authorization header to be leaked when calls to the health
# endpoint are made and fail.
raw_request_headers=self._get_masked_headers(
additional_args.get("headers", {}) or {},
),
error=None,
)
except Exception as e:
self.model_call_details["raw_request_typed_dict"] = (
RawRequestTypedDict(
error=str(e),
)
self.model_call_details[
"raw_request_typed_dict"
] = RawRequestTypedDict(
error=str(e),
)
_metadata["raw_request"] = (
"Unable to Log \
_metadata[
"raw_request"
] = "Unable to Log \
raw request: {}".format(
str(e)
)
str(e)
)
if getattr(self, "logger_fn", None) and callable(self.logger_fn):
try:
@ -1265,13 +1265,13 @@ class Logging(LiteLLMLoggingBaseClass):
for callback in callbacks:
try:
if isinstance(callback, CustomLogger):
response: Optional[MCPPostCallResponseObject] = (
await callback.async_post_mcp_tool_call_hook(
kwargs=kwargs,
response_obj=post_mcp_tool_call_response_obj,
start_time=start_time,
end_time=end_time,
)
response: Optional[
MCPPostCallResponseObject
] = await callback.async_post_mcp_tool_call_hook(
kwargs=kwargs,
response_obj=post_mcp_tool_call_response_obj,
start_time=start_time,
end_time=end_time,
)
######################################################################
# if any of the callbacks modify the response, use the modified response
@ -1466,9 +1466,9 @@ class Logging(LiteLLMLoggingBaseClass):
verbose_logger.debug(
f"response_cost_failure_debug_information: {debug_info}"
)
self.model_call_details["response_cost_failure_debug_information"] = (
debug_info
)
self.model_call_details[
"response_cost_failure_debug_information"
] = debug_info
return None
try:
@ -1494,9 +1494,9 @@ class Logging(LiteLLMLoggingBaseClass):
verbose_logger.debug(
f"response_cost_failure_debug_information: {debug_info}"
)
self.model_call_details["response_cost_failure_debug_information"] = (
debug_info
)
self.model_call_details[
"response_cost_failure_debug_information"
] = debug_info
return None
@ -1652,9 +1652,9 @@ class Logging(LiteLLMLoggingBaseClass):
result=logging_result
)
self.model_call_details["standard_logging_object"] = (
self._build_standard_logging_payload(logging_result, start_time, end_time)
)
self.model_call_details[
"standard_logging_object"
] = self._build_standard_logging_payload(logging_result, start_time, end_time)
if (
standard_logging_payload := self.model_call_details.get(
@ -1732,9 +1732,9 @@ class Logging(LiteLLMLoggingBaseClass):
end_time = datetime.datetime.now()
if self.completion_start_time is None:
self.completion_start_time = end_time
self.model_call_details["completion_start_time"] = (
self.completion_start_time
)
self.model_call_details[
"completion_start_time"
] = self.completion_start_time
self.model_call_details["log_event_type"] = "successful_api_call"
self.model_call_details["end_time"] = end_time
@ -1771,10 +1771,10 @@ class Logging(LiteLLMLoggingBaseClass):
end_time=end_time,
)
elif isinstance(result, dict) or isinstance(result, list):
self.model_call_details["standard_logging_object"] = (
self._build_standard_logging_payload(
result, start_time, end_time
)
self.model_call_details[
"standard_logging_object"
] = self._build_standard_logging_payload(
result, start_time, end_time
)
if (
standard_logging_payload := self.model_call_details.get(
@ -1783,9 +1783,9 @@ class Logging(LiteLLMLoggingBaseClass):
) is not None:
emit_standard_logging_payload(standard_logging_payload)
elif standard_logging_object is not None:
self.model_call_details["standard_logging_object"] = (
standard_logging_object
)
self.model_call_details[
"standard_logging_object"
] = standard_logging_object
else:
self.model_call_details["response_cost"] = None
@ -1943,17 +1943,17 @@ class Logging(LiteLLMLoggingBaseClass):
verbose_logger.debug(
"Logging Details LiteLLM-Success Call streaming complete"
)
self.model_call_details["complete_streaming_response"] = (
complete_streaming_response
)
self.model_call_details["response_cost"] = (
self._response_cost_calculator(result=complete_streaming_response)
)
self.model_call_details[
"complete_streaming_response"
] = complete_streaming_response
self.model_call_details[
"response_cost"
] = self._response_cost_calculator(result=complete_streaming_response)
## STANDARDIZED LOGGING PAYLOAD
self.model_call_details["standard_logging_object"] = (
self._build_standard_logging_payload(
complete_streaming_response, start_time, end_time
)
self.model_call_details[
"standard_logging_object"
] = self._build_standard_logging_payload(
complete_streaming_response, start_time, end_time
)
if (
standard_logging_payload := self.model_call_details.get(
@ -2287,10 +2287,10 @@ class Logging(LiteLLMLoggingBaseClass):
)
else:
if self.stream and complete_streaming_response:
self.model_call_details["complete_response"] = (
self.model_call_details.get(
"complete_streaming_response", {}
)
self.model_call_details[
"complete_response"
] = self.model_call_details.get(
"complete_streaming_response", {}
)
result = self.model_call_details["complete_response"]
openMeterLogger.log_success_event(
@ -2314,10 +2314,10 @@ class Logging(LiteLLMLoggingBaseClass):
)
else:
if self.stream and complete_streaming_response:
self.model_call_details["complete_response"] = (
self.model_call_details.get(
"complete_streaming_response", {}
)
self.model_call_details[
"complete_response"
] = self.model_call_details.get(
"complete_streaming_response", {}
)
result = self.model_call_details["complete_response"]
@ -2456,9 +2456,9 @@ class Logging(LiteLLMLoggingBaseClass):
if complete_streaming_response is not None:
print_verbose("Async success callbacks: Got a complete streaming response")
self.model_call_details["async_complete_streaming_response"] = (
complete_streaming_response
)
self.model_call_details[
"async_complete_streaming_response"
] = complete_streaming_response
try:
if self.model_call_details.get("cache_hit", False) is True:
@ -2469,10 +2469,10 @@ class Logging(LiteLLMLoggingBaseClass):
model_call_details=self.model_call_details
)
# base_model defaults to None if not set on model_info
self.model_call_details["response_cost"] = (
self._response_cost_calculator(
result=complete_streaming_response
)
self.model_call_details[
"response_cost"
] = self._response_cost_calculator(
result=complete_streaming_response
)
verbose_logger.debug(
@ -2485,10 +2485,10 @@ class Logging(LiteLLMLoggingBaseClass):
self.model_call_details["response_cost"] = None
## STANDARDIZED LOGGING PAYLOAD
self.model_call_details["standard_logging_object"] = (
self._build_standard_logging_payload(
complete_streaming_response, start_time, end_time
)
self.model_call_details[
"standard_logging_object"
] = self._build_standard_logging_payload(
complete_streaming_response, start_time, end_time
)
# print standard logging payload
@ -2515,9 +2515,9 @@ class Logging(LiteLLMLoggingBaseClass):
# _success_handler_helper_fn
if self.model_call_details.get("standard_logging_object") is None:
## STANDARDIZED LOGGING PAYLOAD
self.model_call_details["standard_logging_object"] = (
self._build_standard_logging_payload(result, start_time, end_time)
)
self.model_call_details[
"standard_logging_object"
] = self._build_standard_logging_payload(result, start_time, end_time)
# print standard logging payload
if (
@ -2760,18 +2760,18 @@ class Logging(LiteLLMLoggingBaseClass):
## STANDARDIZED LOGGING PAYLOAD
self.model_call_details["standard_logging_object"] = (
get_standard_logging_object_payload(
kwargs=self.model_call_details,
init_response_obj={},
start_time=start_time,
end_time=end_time,
logging_obj=self,
status="failure",
error_str=str(exception),
original_exception=exception,
standard_built_in_tools_params=self.standard_built_in_tools_params,
)
self.model_call_details[
"standard_logging_object"
] = get_standard_logging_object_payload(
kwargs=self.model_call_details,
init_response_obj={},
start_time=start_time,
end_time=end_time,
logging_obj=self,
status="failure",
error_str=str(exception),
original_exception=exception,
standard_built_in_tools_params=self.standard_built_in_tools_params,
)
return start_time, end_time
@ -2920,7 +2920,10 @@ class Logging(LiteLLMLoggingBaseClass):
callback_func=callback,
)
if (
isinstance(callback, CustomLogger) and is_sync_request
isinstance(callback, CustomLogger)
and is_sync_request
and self.call_type
!= CallTypes.pass_through.value # pass-through endpoints call async_log_failure_event
): # custom logger class
callback.log_failure_event(
start_time=start_time,
@ -3735,9 +3738,9 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
service_name=arize_config.project_name,
)
os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = (
f"space_id={arize_config.space_key or arize_config.space_id},api_key={arize_config.api_key}"
)
os.environ[
"OTEL_EXPORTER_OTLP_TRACES_HEADERS"
] = f"space_id={arize_config.space_key or arize_config.space_id},api_key={arize_config.api_key}"
for callback in _in_memory_loggers:
if (
isinstance(callback, ArizeLogger)
@ -3763,13 +3766,13 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
existing_attrs = os.environ.get("OTEL_RESOURCE_ATTRIBUTES", "")
# Add openinference.project.name attribute
if existing_attrs:
os.environ["OTEL_RESOURCE_ATTRIBUTES"] = (
f"{existing_attrs},openinference.project.name={arize_phoenix_config.project_name}"
)
os.environ[
"OTEL_RESOURCE_ATTRIBUTES"
] = f"{existing_attrs},openinference.project.name={arize_phoenix_config.project_name}"
else:
os.environ["OTEL_RESOURCE_ATTRIBUTES"] = (
f"openinference.project.name={arize_phoenix_config.project_name}"
)
os.environ[
"OTEL_RESOURCE_ATTRIBUTES"
] = f"openinference.project.name={arize_phoenix_config.project_name}"
# Set Phoenix project name from environment variable
phoenix_project_name = os.environ.get("PHOENIX_PROJECT_NAME", None)
@ -3777,19 +3780,19 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
existing_attrs = os.environ.get("OTEL_RESOURCE_ATTRIBUTES", "")
# Add openinference.project.name attribute
if existing_attrs:
os.environ["OTEL_RESOURCE_ATTRIBUTES"] = (
f"{existing_attrs},openinference.project.name={phoenix_project_name}"
)
os.environ[
"OTEL_RESOURCE_ATTRIBUTES"
] = f"{existing_attrs},openinference.project.name={phoenix_project_name}"
else:
os.environ["OTEL_RESOURCE_ATTRIBUTES"] = (
f"openinference.project.name={phoenix_project_name}"
)
os.environ[
"OTEL_RESOURCE_ATTRIBUTES"
] = f"openinference.project.name={phoenix_project_name}"
# auth can be disabled on local deployments of arize phoenix
if arize_phoenix_config.otlp_auth_headers is not None:
os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = (
arize_phoenix_config.otlp_auth_headers
)
os.environ[
"OTEL_EXPORTER_OTLP_TRACES_HEADERS"
] = arize_phoenix_config.otlp_auth_headers
for callback in _in_memory_loggers:
if (
@ -3965,9 +3968,9 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
exporter="otlp_http",
endpoint="https://langtrace.ai/api/trace",
)
os.environ["OTEL_EXPORTER_OTLP_TRACES_HEADERS"] = (
f"api_key={os.getenv('LANGTRACE_API_KEY')}"
)
os.environ[
"OTEL_EXPORTER_OTLP_TRACES_HEADERS"
] = f"api_key={os.getenv('LANGTRACE_API_KEY')}"
for callback in _in_memory_loggers:
if (
isinstance(callback, OpenTelemetry)
@ -4881,10 +4884,10 @@ class StandardLoggingPayloadSetup:
for key in StandardLoggingHiddenParams.__annotations__.keys():
if key in hidden_params:
if key == "additional_headers":
clean_hidden_params["additional_headers"] = (
StandardLoggingPayloadSetup.get_additional_headers(
hidden_params[key]
)
clean_hidden_params[
"additional_headers"
] = StandardLoggingPayloadSetup.get_additional_headers(
hidden_params[key]
)
else:
clean_hidden_params[key] = hidden_params[key] # type: ignore
@ -5036,7 +5039,6 @@ class StandardLoggingPayloadSetup:
dynamic_litellm_session_id = litellm_params.get("litellm_session_id")
dynamic_litellm_trace_id = litellm_params.get("litellm_trace_id")
# Note: we recommend using `litellm_session_id` for session tracking
# `litellm_trace_id` is an internal litellm param
if dynamic_litellm_session_id:
@ -5507,9 +5509,9 @@ def scrub_sensitive_keys_in_metadata(litellm_params: Optional[dict]):
):
for k, v in metadata["user_api_key_metadata"].items():
if k == "logging": # prevent logging user logging keys
cleaned_user_api_key_metadata[k] = (
"scrubbed_by_litellm_for_sensitive_keys"
)
cleaned_user_api_key_metadata[
k
] = "scrubbed_by_litellm_for_sensitive_keys"
else:
cleaned_user_api_key_metadata[k] = v

View file

@ -957,6 +957,13 @@ async def pass_through_request( # noqa: PLR0915
if "custom_llm_provider" not in request_payload and custom_llm_provider:
request_payload["custom_llm_provider"] = custom_llm_provider
# Pass the existing logging_obj so _handle_logging_proxy_only_error
# uses it (preserves call_type="pass_through_endpoint" for dedup)
try:
request_payload["litellm_logging_obj"] = logging_obj
except NameError:
pass
await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict,
original_exception=e,
@ -1697,6 +1704,12 @@ async def websocket_passthrough_request( # noqa: PLR0915
for key, value in kwargs.items():
request_payload[key] = value
# Pass the existing logging_obj (preserves call_type for dedup)
try:
request_payload["litellm_logging_obj"] = logging_obj
except NameError:
pass
# Log the connection failure using the same pattern as HTTP
await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict,
@ -1723,6 +1736,12 @@ async def websocket_passthrough_request( # noqa: PLR0915
for key, value in kwargs.items():
request_payload[key] = value
# Pass the existing logging_obj (preserves call_type for dedup)
try:
request_payload["litellm_logging_obj"] = logging_obj
except NameError:
pass
# Log the unexpected error using the same pattern as HTTP
await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict,
@ -2114,11 +2133,7 @@ class InitPassThroughEndpointHelpers:
# If path matches and method filter is provided, check if method is allowed
if path_matches:
if (
method is None
or not route_methods
or method in route_methods
):
if method is None or not route_methods or method in route_methods:
return _registered_pass_through_routes[key]
return None

View file

@ -255,6 +255,96 @@ async def test_pass_through_request_failure_handler():
assert "traceback_str" in call_args
def test_pass_through_failure_no_duplicate_custom_logger_callbacks():
"""
Test that pass-through endpoint failures don't trigger duplicate CustomLogger callbacks.
The sync failure_handler should skip CustomLogger.log_failure_event() when
call_type is "pass_through_endpoint", since async_failure_handler already
calls async_log_failure_event(). This prevents duplicate logs to Datadog/Arize.
"""
from unittest.mock import MagicMock
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging
# Create a mock CustomLogger
mock_callback = MagicMock(spec=CustomLogger)
mock_callback.log_failure_event = MagicMock()
# Create a logging object with pass_through call_type
logging_obj = Logging(
model="unknown",
messages=[],
stream=False,
call_type="pass_through_endpoint",
start_time=None,
litellm_call_id="test-call-id",
function_id="test",
)
import litellm
original_failure_callback = litellm.failure_callback[:]
try:
# Register the mock callback in the sync failure list
litellm.failure_callback = [mock_callback]
# Call the sync failure_handler
logging_obj.failure_handler(
exception=Exception("test error"),
traceback_exception="test traceback",
)
# log_failure_event should NOT be called for pass-through endpoints
mock_callback.log_failure_event.assert_not_called()
finally:
litellm.failure_callback = original_failure_callback
def test_non_pass_through_failure_calls_custom_logger_callback():
"""
Test that non-pass-through endpoint failures still call CustomLogger.log_failure_event().
Ensures the pass-through guard doesn't break normal failure logging.
"""
from unittest.mock import MagicMock
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging
# Create a mock CustomLogger
mock_callback = MagicMock(spec=CustomLogger)
mock_callback.log_failure_event = MagicMock()
# Create a logging object with a NON-pass-through call_type
logging_obj = Logging(
model="test-model",
messages=[],
stream=False,
call_type="completion",
start_time=None,
litellm_call_id="test-call-id",
function_id="test",
)
import litellm
original_failure_callback = litellm.failure_callback[:]
try:
litellm.failure_callback = [mock_callback]
logging_obj.failure_handler(
exception=Exception("test error"),
traceback_exception="test traceback",
)
# log_failure_event SHOULD be called for non-pass-through endpoints
mock_callback.log_failure_event.assert_called_once()
finally:
litellm.failure_callback = original_failure_callback
def test_is_langfuse_route():
"""
Test that the is_langfuse_route method correctly identifies Langfuse routes
@ -2232,11 +2322,15 @@ def test_build_full_path_with_root_default():
InitPassThroughEndpointHelpers,
)
with patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_server_root_path") as mock_get_root:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_server_root_path"
) as mock_get_root:
# Test with default root path
mock_get_root.return_value = "/"
result = InitPassThroughEndpointHelpers._build_full_path_with_root("/api/v1/endpoint")
result = InitPassThroughEndpointHelpers._build_full_path_with_root(
"/api/v1/endpoint"
)
assert result == "/api/v1/endpoint"
@ -2248,11 +2342,15 @@ def test_build_full_path_with_root_custom():
InitPassThroughEndpointHelpers,
)
with patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_server_root_path") as mock_get_root:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_server_root_path"
) as mock_get_root:
# Test with custom root path /proxy
mock_get_root.return_value = "/proxy"
result = InitPassThroughEndpointHelpers._build_full_path_with_root("/api/v1/endpoint")
result = InitPassThroughEndpointHelpers._build_full_path_with_root(
"/api/v1/endpoint"
)
assert result == "/proxy/api/v1/endpoint"
@ -2264,7 +2362,9 @@ def test_build_full_path_with_root_nested():
InitPassThroughEndpointHelpers,
)
with patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_server_root_path") as mock_get_root:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_server_root_path"
) as mock_get_root:
# Test with nested root path /api/v2
mock_get_root.return_value = "/api/v2"
@ -2296,24 +2396,46 @@ def test_is_registered_pass_through_route_with_custom_root():
"headers": {},
}
with patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_server_root_path") as mock_get_root:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_server_root_path"
) as mock_get_root:
# Test with custom root path /proxy
mock_get_root.return_value = "/proxy"
# Should match when request route includes the root path
assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/proxy/api/endpoint") is True
assert (
InitPassThroughEndpointHelpers.is_registered_pass_through_route(
"/proxy/api/endpoint"
)
is True
)
# Should not match when request route doesn't include root path
assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/api/endpoint") is False
assert (
InitPassThroughEndpointHelpers.is_registered_pass_through_route(
"/api/endpoint"
)
is False
)
# Test with default root path
mock_get_root.return_value = "/"
# Should match with default root
assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/api/endpoint") is True
assert (
InitPassThroughEndpointHelpers.is_registered_pass_through_route(
"/api/endpoint"
)
is True
)
# Should not match with root prepended when root is /
assert InitPassThroughEndpointHelpers.is_registered_pass_through_route("/proxy/api/endpoint") is False
assert (
InitPassThroughEndpointHelpers.is_registered_pass_through_route(
"/proxy/api/endpoint"
)
is False
)
# Clean up
_registered_pass_through_routes.clear()
@ -2345,25 +2467,33 @@ def test_get_registered_pass_through_route_with_custom_root():
route_key = f"{endpoint_id}:exact:{path}"
_registered_pass_through_routes[route_key] = target_config
with patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_server_root_path") as mock_get_root:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_server_root_path"
) as mock_get_root:
# Test with custom root path /litellm
mock_get_root.return_value = "/litellm"
# Should return config when request route includes root path
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"
# Should return None when route doesn't match
result = InitPassThroughEndpointHelpers.get_registered_pass_through_route("/chat/completions")
result = InitPassThroughEndpointHelpers.get_registered_pass_through_route(
"/chat/completions"
)
assert result is None
# Test with default root path
mock_get_root.return_value = "/"
# Should return config with default root
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"