mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
style(bedrock/claude_platform): fix ruff format + add coverage test
Run ruff format on common_utils.py. Add test for the invalid-type fallback branch in resolve_unsupported_override (line 95 coverage).
This commit is contained in:
parent
2e55105020
commit
40e44ed4d6
2 changed files with 33 additions and 19 deletions
|
|
@ -3,15 +3,14 @@ from typing import Literal, Optional, Protocol, Tuple
|
|||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
|
||||
class _SupportsGet(Protocol):
|
||||
def get(self, key: str, default: object = None) -> object: ...
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
CLAUDE_PLATFORM_SERVICE_NAME: Literal["aws-external-anthropic"] = (
|
||||
"aws-external-anthropic"
|
||||
)
|
||||
|
||||
CLAUDE_PLATFORM_SERVICE_NAME: Literal["aws-external-anthropic"] = "aws-external-anthropic"
|
||||
CLAUDE_PLATFORM_BEDROCK_ROUTE = "claude_platform/"
|
||||
|
||||
# Auth/routing params consumed by validate_environment / sign_request that
|
||||
|
|
@ -56,9 +55,7 @@ def filter_claude_platform_request_body(
|
|||
Filters a copy so callers' ``sign_request`` still sees ``aws_region_name``.
|
||||
"""
|
||||
unsupported = (
|
||||
unsupported_override
|
||||
if unsupported_override is not None
|
||||
else CLAUDE_PLATFORM_ON_AWS_UNSUPPORTED_REQUEST_PARAMS
|
||||
unsupported_override if unsupported_override is not None else CLAUDE_PLATFORM_ON_AWS_UNSUPPORTED_REQUEST_PARAMS
|
||||
)
|
||||
dropped_unsupported = [k for k in params if k in unsupported]
|
||||
if dropped_unsupported:
|
||||
|
|
@ -72,9 +69,7 @@ def filter_claude_platform_request_body(
|
|||
return {
|
||||
k: v
|
||||
for k, v in params.items()
|
||||
if k not in CLAUDE_PLATFORM_ON_AWS_NON_REQUEST_PARAMS
|
||||
and k not in unsupported
|
||||
and not k.startswith("aws_")
|
||||
if k not in CLAUDE_PLATFORM_ON_AWS_NON_REQUEST_PARAMS and k not in unsupported and not k.startswith("aws_")
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -113,14 +108,10 @@ class BedrockClaudePlatformMixin(BaseAWSLLM):
|
|||
or litellm_params.get("anthropic-workspace-id")
|
||||
)
|
||||
if workspace_id is None:
|
||||
workspace_id = optional_params.get(
|
||||
"anthropic_workspace_id"
|
||||
) or litellm_params.get("anthropic_workspace_id")
|
||||
workspace_id = optional_params.get("anthropic_workspace_id") or litellm_params.get("anthropic_workspace_id")
|
||||
if workspace_id is not None:
|
||||
return str(workspace_id)
|
||||
return get_secret_str("ANTHROPIC_AWS_WORKSPACE_ID") or get_secret_str(
|
||||
"ANTHROPIC_WORKSPACE_ID"
|
||||
)
|
||||
return get_secret_str("ANTHROPIC_AWS_WORKSPACE_ID") or get_secret_str("ANTHROPIC_WORKSPACE_ID")
|
||||
|
||||
def _get_required_aws_region_name(self, optional_params: dict) -> str:
|
||||
aws_region_name = (
|
||||
|
|
@ -158,9 +149,7 @@ class BedrockClaudePlatformMixin(BaseAWSLLM):
|
|||
)
|
||||
if api_base is None:
|
||||
aws_region_name = self._get_required_aws_region_name(optional_params)
|
||||
api_base = (
|
||||
f"https://{CLAUDE_PLATFORM_SERVICE_NAME}.{aws_region_name}.api.aws"
|
||||
)
|
||||
api_base = f"https://{CLAUDE_PLATFORM_SERVICE_NAME}.{aws_region_name}.api.aws"
|
||||
if not api_base.endswith("/v1/messages"):
|
||||
api_base = f"{api_base.rstrip('/')}/v1/messages"
|
||||
return api_base
|
||||
|
|
|
|||
|
|
@ -480,6 +480,31 @@ def test_claude_platform_unsupported_override_allows_context_management():
|
|||
assert "context_management" in request_body
|
||||
|
||||
|
||||
def test_claude_platform_unsupported_override_ignores_invalid_type():
|
||||
"""
|
||||
If claude_platform_unsupported_params is set to a non-collection type
|
||||
(e.g. a string), the override is ignored and defaults apply.
|
||||
"""
|
||||
from litellm.llms.bedrock.claude_platform.transformation import (
|
||||
BedrockClaudePlatformConfig,
|
||||
)
|
||||
|
||||
config = BedrockClaudePlatformConfig()
|
||||
request_body = config.transform_request(
|
||||
model="claude-sonnet-4-6",
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
optional_params={
|
||||
"context_management": {"edits": [{"type": "clear_tool_uses_20250919"}]},
|
||||
"max_tokens": 10,
|
||||
},
|
||||
litellm_params={"claude_platform_unsupported_params": "not_a_list"},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert "context_management" not in request_body
|
||||
assert request_body["max_tokens"] == 10
|
||||
|
||||
|
||||
def test_chat_completion_claude_platform_sigv4_body_has_no_auth_params():
|
||||
"""
|
||||
End-to-end (mocked transport): a config-driven SigV4 call with
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue