mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
commit
230db7e161
2 changed files with 77 additions and 0 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue