mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: scope Bedrock bearer keys
This commit is contained in:
parent
31a529aee6
commit
53c90c351b
2 changed files with 97 additions and 10 deletions
|
|
@ -677,25 +677,72 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
except Exception:
|
||||
return model.partition("/")[0]
|
||||
|
||||
@staticmethod
|
||||
def _router_allows_bedrock(request_data: Mapping[str, object]) -> bool | None:
|
||||
model: Final[object | None] = request_data.get("model")
|
||||
if not isinstance(model, str):
|
||||
return False
|
||||
|
||||
try:
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
except ImportError:
|
||||
return None
|
||||
if llm_router is None:
|
||||
return None
|
||||
|
||||
metadata: Final = request_data.get("metadata")
|
||||
litellm_metadata: Final = request_data.get("litellm_metadata")
|
||||
team_id: Final[object | None] = (
|
||||
metadata.get("user_api_key_team_id")
|
||||
if isinstance(metadata, Mapping)
|
||||
else litellm_metadata.get("user_api_key_team_id")
|
||||
if isinstance(litellm_metadata, Mapping)
|
||||
else None
|
||||
)
|
||||
try:
|
||||
deployments: Final = llm_router.get_model_list(
|
||||
model_name=model,
|
||||
team_id=team_id if isinstance(team_id, str) else None,
|
||||
) or []
|
||||
except Exception:
|
||||
return False
|
||||
if not deployments:
|
||||
return False
|
||||
|
||||
providers: list[str] = []
|
||||
for deployment in deployments:
|
||||
params: object = deployment.get("litellm_params") if isinstance(deployment, Mapping) else None
|
||||
provider: object = params.get("custom_llm_provider") if isinstance(params, Mapping) else None
|
||||
if not isinstance(provider, str) and isinstance(params, Mapping):
|
||||
deployment_model: Final[object | None] = params.get("model")
|
||||
provider = (
|
||||
BedrockGuardrail._resolve_model_provider(deployment_model)
|
||||
if isinstance(deployment_model, str)
|
||||
else None
|
||||
)
|
||||
if not isinstance(provider, str):
|
||||
return False
|
||||
providers.append(provider)
|
||||
return bool(providers) and all(provider in ("bedrock", "bedrock_converse") for provider in providers)
|
||||
|
||||
@staticmethod
|
||||
def _get_bedrock_api_key(request_data: Mapping[str, object] | None) -> str | None:
|
||||
if not request_data:
|
||||
return None
|
||||
|
||||
top_level_provider: Final[object | None] = request_data.get("custom_llm_provider")
|
||||
request_provider: Final[str | None] = (
|
||||
top_level_provider if isinstance(top_level_provider, str) else None
|
||||
)
|
||||
api_key: Final[object | None] = request_data.get("api_key")
|
||||
if not isinstance(api_key, str):
|
||||
return None
|
||||
|
||||
router_allows_bedrock: Final = BedrockGuardrail._router_allows_bedrock(request_data)
|
||||
if router_allows_bedrock is not None:
|
||||
return api_key if router_allows_bedrock else None
|
||||
|
||||
model: Final[object | None] = request_data.get("model")
|
||||
model_provider: Final[str | None] = (
|
||||
BedrockGuardrail._resolve_model_provider(model) if isinstance(model, str) else None
|
||||
)
|
||||
custom_llm_provider: Final[str | None] = request_provider or model_provider
|
||||
if custom_llm_provider not in ("bedrock", "bedrock_converse"):
|
||||
return None
|
||||
|
||||
api_key: Final[object | None] = request_data.get("api_key")
|
||||
return api_key if isinstance(api_key, str) else None
|
||||
return api_key if model_provider in ("bedrock", "bedrock_converse") else None
|
||||
|
||||
def _load_credentials(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -1167,6 +1167,46 @@ async def test_make_apply_guardrail_request_scopes_api_key_to_bedrock_provider(
|
|||
)
|
||||
|
||||
|
||||
def test_bedrock_api_key_rejects_caller_provider_spoofing(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
{"litellm_params": {"model": "gpt-4o", "custom_llm_provider": "openai"}}
|
||||
]
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert (
|
||||
BedrockGuardrail._get_bedrock_api_key(
|
||||
{
|
||||
"model": "shared-alias",
|
||||
"custom_llm_provider": "bedrock",
|
||||
"api_key": "bedrock-key",
|
||||
}
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_bedrock_api_key_accepts_alias_with_only_bedrock_deployments(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = [
|
||||
{"litellm_params": {"model": "amazon.nova-lite-v1:0", "custom_llm_provider": "bedrock"}}
|
||||
]
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
assert (
|
||||
BedrockGuardrail._get_bedrock_api_key(
|
||||
{"model": "bedrock-alias", "api_key": "bedrock-key"}
|
||||
)
|
||||
== "bedrock-key"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_apply_guardrail_response_uses_OUTPUT_source():
|
||||
"""input_type='response' must call Bedrock with source=OUTPUT and assistant content.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue