[Bug Fix]: Hooks broken on /bedrock passthrough due to missing metadata (#15849)

* refactor handle_bedrock_passthrough_router_model

* test_bedrock_router_passthrough_metadata_initialization
This commit is contained in:
Ishaan Jaff 2025-10-23 11:52:37 -07:00 • committed by GitHub
parent 8c5118195d
commit 511d435f6f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 171 additions and 75 deletions

View file

@ -545,12 +545,23 @@ async def handle_bedrock_passthrough_router_model(
request: Request,
request_body: dict,
llm_router: litellm.Router,
user_api_key_dict: UserAPIKeyAuth,
proxy_logging_obj,
general_settings: dict,
proxy_config,
select_data_generator,
user_model: Optional[str],
user_temperature: Optional[float],
user_request_timeout: Optional[float],
user_max_tokens: Optional[int],
user_api_base: Optional[str],
version: Optional[str],
) -> Union[Response, StreamingResponse]:
"""
Handle Bedrock passthrough for router models (models defined in config.yaml).
This helper delegates to llm_router.allm_passthrough_route for proper credential
and configuration management from the router.
Uses the same common processing path as non-router models to ensure
metadata and hooks are properly initialized.
Args:
model: The router model name (e.g., "aws/anthropic/bedrock-claude-3-5-sonnet-v1")
@ -558,10 +569,16 @@ async def handle_bedrock_passthrough_router_model(
request: The FastAPI request object
request_body: The parsed request body
llm_router: The LiteLLM router instance
user_api_key_dict: The user API key authentication dictionary
(additional args for common processing)
Returns:
Response or StreamingResponse depending on endpoint type
"""
from fastapi import Response as FastAPIResponse
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
# Detect streaming based on endpoint
is_streaming = any(action in endpoint for action in BEDROCK_STREAMING_ACTIONS)
@ -569,83 +586,45 @@ async def handle_bedrock_passthrough_router_model(
f"Bedrock router passthrough: model='{model}', endpoint='{endpoint}', streaming={is_streaming}"
)
# Call router passthrough
# Use the common processing path (same as non-router models)
# This ensures all metadata, hooks, and logging are properly initialized
data: Dict[str, Any] = {}
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
data["model"] = model
data["method"] = request.method
data["endpoint"] = endpoint
data["data"] = request_body
data["custom_llm_provider"] = "bedrock"
# Use the common passthrough processing to handle metadata and hooks
# This also handles all response formatting (streaming/non-streaming) and exceptions
try:
result = await llm_router.allm_passthrough_route(
result = await base_llm_response_processor.base_passthrough_process_llm_request(
request=request,
fastapi_response=FastAPIResponse(),
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
general_settings=general_settings,
proxy_config=proxy_config,
select_data_generator=select_data_generator,
model=model,
method=request.method,
endpoint=endpoint,
request_query_params=request.query_params,
request_headers=dict(request.headers),
stream=is_streaming,
content=None,
data=None,
files=None,
json=(
request_body
if request.headers.get("content-type") == "application/json"
else None
),
params=None,
headers=None,
cookies=None,
)
except httpx.HTTPStatusError as e:
# Handle HTTP errors from the provider by converting to HTTPException
error_body = await e.response.aread()
error_text = error_body.decode("utf-8")
raise HTTPException(
status_code=e.response.status_code,
detail={"error": error_text},
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
version=version,
)
return result
except Exception as e:
from litellm.llms.base_llm.chat.transformation import BaseLLMException
# If it's a BaseLLMException (from non-HTTP errors), convert to HTTPException
if isinstance(e, BaseLLMException):
raise HTTPException(
status_code=e.status_code,
detail={"error": e.message},
)
# Re-raise any other exceptions
raise e
# Handle streaming response
if is_streaming:
import inspect
if inspect.isasyncgen(result):
# AsyncGenerator case
return StreamingResponse(
content=result,
status_code=200,
headers={"content-type": "application/vnd.amazon.eventstream"},
)
else:
# httpx.Response case
result = cast(httpx.Response, result)
return StreamingResponse(
content=result.aiter_bytes(),
status_code=result.status_code,
headers=HttpPassThroughEndpointHelpers.get_response_headers(
headers=result.headers,
custom_headers=None,
),
)
# Handle non-streaming response
result = cast(httpx.Response, result)
content = await result.aread()
return Response(
content=content,
status_code=result.status_code,
headers=HttpPassThroughEndpointHelpers.get_response_headers(
headers=result.headers,
custom_headers=None,
),
)
# Use common exception handling
raise await base_llm_response_processor._handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
)
async def handle_bedrock_count_tokens(
@ -778,6 +757,7 @@ async def bedrock_llm_proxy_route(
)
# If router model, use dedicated router passthrough handler
# This uses the same common processing path as non-router models
if is_router_model and llm_router:
return await handle_bedrock_passthrough_router_model(
model=model,
@ -785,6 +765,17 @@ async def bedrock_llm_proxy_route(
request=request,
request_body=request_body,
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
general_settings=general_settings,
proxy_config=proxy_config,
select_data_generator=select_data_generator,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
version=version,
)
# Fall back to existing implementation for direct Bedrock models

View file

@ -1693,3 +1693,108 @@ async def test_filter_endpoints_by_team_allowed_routes_partial_match():
assert len(result) == 2
assert result[0].path == "/api/openai"
assert result[1].path == "/api/azure"
@pytest.mark.asyncio
async def test_bedrock_router_passthrough_metadata_initialization():
"""
Test that bedrock router passthrough properly initializes metadata for hooks.
This test verifies the fix for issue #15826 where metadata.headers and
litellm_params.proxy_server_request were missing for /bedrock passthrough
requests with router models.
The fix ensures router bedrock models use the same common processing path
as non-router models, which properly initializes all metadata structures.
"""
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
handle_bedrock_passthrough_router_model,
)
# Mock ProxyBaseLLMRequestProcessing to verify it's used
with patch(
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing"
) as mock_processing_class:
# Setup mock instance
mock_processor = MagicMock()
mock_processing_class.return_value = mock_processor
# Mock successful response
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.aread = AsyncMock(return_value=b'{"content": [{"text": "Hello"}]}')
mock_processor.base_passthrough_process_llm_request = AsyncMock(
return_value=mock_response
)
# Create mock request with headers
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://localhost:4000/bedrock/model/my-model/invoke"
mock_request.headers = Headers({
"content-type": "application/json",
"authorization": "Bearer sk-test-key",
"x-custom-header": "test-value"
})
mock_request.query_params = QueryParams({})
# Create mock user API key dict with all required fields
mock_user_api_key_dict = MagicMock()
mock_user_api_key_dict.api_key = "sk-test-key"
mock_user_api_key_dict.key_alias = "test-alias"
mock_user_api_key_dict.user_id = "user-123"
mock_user_api_key_dict.team_id = "team-123"
# Mock other required dependencies
mock_router = MagicMock()
mock_proxy_logging = MagicMock()
mock_general_settings = {}
mock_proxy_config = MagicMock()
mock_select_data_generator = MagicMock()
request_body = {
"max_tokens": 100,
"messages": [{"role": "user", "content": "Hello"}],
"anthropic_version": "bedrock-2023-05-31"
}
# Call the function
result = await handle_bedrock_passthrough_router_model(
model="my-bedrock-model",
endpoint="/model/my-bedrock-model/invoke",
request=mock_request,
request_body=request_body,
llm_router=mock_router,
user_api_key_dict=mock_user_api_key_dict,
proxy_logging_obj=mock_proxy_logging,
general_settings=mock_general_settings,
proxy_config=mock_proxy_config,
select_data_generator=mock_select_data_generator,
user_model=None,
user_temperature=None,
user_request_timeout=None,
user_max_tokens=None,
user_api_base=None,
version="1.0",
)
# Verify that ProxyBaseLLMRequestProcessing was instantiated
# This is the KEY assertion - router models now use the common processing path
mock_processing_class.assert_called_once()
# Verify that base_passthrough_process_llm_request was called
# This proves we're using the common processing path that initializes metadata
mock_processor.base_passthrough_process_llm_request.assert_called_once()
# Verify the call included all required parameters for proper metadata initialization
call_kwargs = mock_processor.base_passthrough_process_llm_request.call_args[1]
# These are the critical parameters that ensure metadata is properly initialized:
assert call_kwargs["request"] == mock_request, "Request must be passed for header extraction"
assert call_kwargs["user_api_key_dict"] == mock_user_api_key_dict, "User API key dict needed for metadata"
assert call_kwargs["proxy_logging_obj"] == mock_proxy_logging, "Logging obj needed for hooks"
assert call_kwargs["llm_router"] == mock_router, "Router needed for model routing"
assert call_kwargs["model"] == "my-bedrock-model", "Model name must be passed"
# Verify response was returned
assert result == mock_response