mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
[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:
parent
8c5118195d
commit
511d435f6f
2 changed files with 171 additions and 75 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue