Merge pull request #18100 from BerriAI/litellm_bedrock_qwen_arn_fix

fix: Add qwen 2 and qwen 3 in get_bedrock_model_id
This commit is contained in:
Sameer Kankute 2025-12-17 22:34:05 +05:30 • committed by GitHub
commit 230db7e161
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 77 additions and 0 deletions

View file

@ -357,6 +357,14 @@ class BaseAWSLLM:
model_id = BaseAWSLLM._get_model_id_from_model_with_spec(
model_id, spec="openai"
)
elif provider == "qwen2" and "qwen2/" in model_id:
model_id = BaseAWSLLM._get_model_id_from_model_with_spec(
model_id, spec="qwen2"
)
elif provider == "qwen3" and "qwen3/" in model_id:
model_id = BaseAWSLLM._get_model_id_from_model_with_spec(
model_id, spec="qwen3"
)
return model_id
@staticmethod

View file

@ -290,3 +290,72 @@ def test_qwen2_provider_detection():
assert config is not None
assert isinstance(config, AmazonQwen2Config)
def test_qwen2_model_id_extraction_with_arn():
"""Test that model ID is correctly extracted from bedrock/qwen2/arn... paths"""
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
# Test case: bedrock/qwen2/arn:aws:bedrock:us-east-1:123456789012:imported-model/test-qwen2
# The qwen2/ prefix should be stripped, leaving only the ARN for encoding
model = "qwen2/arn:aws:bedrock:us-east-1:123456789012:imported-model/test-qwen2"
provider = "qwen2"
result = BaseAWSLLM.get_bedrock_model_id(
optional_params={},
provider=provider,
model=model
)
# The result should NOT contain "qwen2/" - it should be stripped
assert "qwen2/" not in result
# The result should be URL-encoded ARN
assert "arn%3Aaws%3Abedrock" in result or "arn:aws:bedrock" in result
def test_qwen2_model_id_extraction_without_qwen2_prefix():
"""Test that model ID extraction doesn't strip qwen2/ when provider is not qwen2"""
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
# Test case: just a model name without qwen2/ prefix
model = "arn:aws:bedrock:us-east-1:123456789012:imported-model/test-qwen2"
provider = "qwen2"
result = BaseAWSLLM.get_bedrock_model_id(
optional_params={},
provider=provider,
model=model
)
# Result should be encoded ARN
assert "arn" in result.lower() or "aws" in result.lower()
def test_qwen2_get_bedrock_model_id_with_various_formats():
"""Test get_bedrock_model_id with various Qwen2 model path formats"""
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
test_cases = [
{
"model": "qwen2/arn:aws:bedrock:us-east-1:123456789012:imported-model/test-qwen2",
"provider": "qwen2",
"should_not_contain": "qwen2/",
"description": "Qwen2 imported model ARN"
},
{
"model": "bedrock/qwen2/arn:aws:bedrock:us-east-1:123456789012:imported-model/test-qwen2",
"provider": "qwen2",
"should_not_contain": "qwen2/",
"description": "Bedrock prefixed Qwen2 ARN"
}
]
for test_case in test_cases:
result = BaseAWSLLM.get_bedrock_model_id(
optional_params={},
provider=test_case["provider"],
model=test_case["model"]
)
assert test_case["should_not_contain"] not in result, \
f"Failed for {test_case['description']}: {test_case['should_not_contain']} found in {result}"