mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix: apply custom video pricing from deployment model_info (#21923)
* auth_with_role_name add region_name arg for cross-account sts * update tests to include case with aws_region_name for _auth_with_aws_role * Only pass region_name to STS client when aws_region_name is set * Add optional aws_sts_endpoint to _auth_with_aws_role * Parametrize ambient-credentials test for no opts, region_name, and aws_sts_endpoint * consistently passing region and endpoint args into explicit credentials irsa * fix env var leakage * fix: bedrock openai-compatible imported-model should also have model arn encoded * fix: custom pricing not applied for /v1/videos endpoint (#21907) * fix: resolve mypy type errors for video pricing model_info parameter Use Optional[ModelInfo] instead of Optional[dict] and restructure cost_info narrowing so mypy can properly track non-None state. --------- Co-authored-by: An Tang <ta@stripe.com> Co-authored-by: Sameer Kankute <sameer@berri.ai>
This commit is contained in:
parent
058852ac5b
commit
289e6031bc
7 changed files with 276 additions and 133 deletions
|
|
@ -1195,6 +1195,16 @@ def completion_cost( # noqa: PLR0915
|
|||
)
|
||||
elif call_type in _VIDEO_CALL_TYPES:
|
||||
### VIDEO GENERATION COST CALCULATION ###
|
||||
# Extract custom model_info for deployment-specific pricing
|
||||
_video_model_info: Optional[ModelInfo] = None
|
||||
if custom_pricing and litellm_logging_obj is not None:
|
||||
_litellm_params = getattr(
|
||||
litellm_logging_obj, "litellm_params", None
|
||||
)
|
||||
if _litellm_params is not None:
|
||||
_metadata = _litellm_params.get("metadata", {}) or {}
|
||||
_video_model_info = _metadata.get("model_info", None)
|
||||
|
||||
usage_obj = getattr(completion_response, "usage", None)
|
||||
if completion_response is not None and usage_obj:
|
||||
# Handle both dict and Pydantic Usage object
|
||||
|
|
@ -1215,12 +1225,14 @@ def completion_cost( # noqa: PLR0915
|
|||
model=model,
|
||||
duration_seconds=duration_seconds,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model_info=_video_model_info,
|
||||
)
|
||||
# Fallback to default video cost calculation if no duration available
|
||||
return default_video_cost_calculator(
|
||||
model=model,
|
||||
duration_seconds=0.0, # Default to 0 if no duration available
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model_info=_video_model_info,
|
||||
)
|
||||
elif call_type in _SPEECH_CALL_TYPES:
|
||||
prompt_characters = litellm.utils._count_characters(text=prompt)
|
||||
|
|
@ -1845,6 +1857,7 @@ def default_video_cost_calculator(
|
|||
model: str,
|
||||
duration_seconds: float,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
model_info: Optional[ModelInfo] = None,
|
||||
) -> float:
|
||||
"""
|
||||
Default video cost calculator for video generation
|
||||
|
|
@ -1853,6 +1866,9 @@ def default_video_cost_calculator(
|
|||
model (str): Model name
|
||||
duration_seconds (float): Duration of the generated video in seconds
|
||||
custom_llm_provider (Optional[str]): Custom LLM provider
|
||||
model_info (Optional[ModelInfo]): Deployment-level model info containing
|
||||
custom video pricing. When provided, used before falling back to
|
||||
the global litellm.model_cost lookup.
|
||||
|
||||
Returns:
|
||||
float: Cost in USD for the video generation
|
||||
|
|
@ -1860,42 +1876,47 @@ def default_video_cost_calculator(
|
|||
Raises:
|
||||
Exception: If model pricing not found in cost map
|
||||
"""
|
||||
# Build model names for cost lookup
|
||||
base_model_name = model
|
||||
model_name_without_custom_llm_provider: Optional[str] = None
|
||||
if custom_llm_provider and model.startswith(f"{custom_llm_provider}/"):
|
||||
model_name_without_custom_llm_provider = model.replace(
|
||||
f"{custom_llm_provider}/", ""
|
||||
)
|
||||
base_model_name = (
|
||||
f"{custom_llm_provider}/{model_name_without_custom_llm_provider}"
|
||||
)
|
||||
|
||||
verbose_logger.debug(f"Looking up cost for video model: {base_model_name}")
|
||||
|
||||
model_without_provider = model.split("/")[-1]
|
||||
|
||||
# Try model with provider first, fall back to base model name
|
||||
# Use custom model_info pricing if provided (deployment-specific pricing)
|
||||
cost_info: Optional[dict] = None
|
||||
models_to_check: List[Optional[str]] = [
|
||||
base_model_name,
|
||||
model,
|
||||
model_without_provider,
|
||||
model_name_without_custom_llm_provider,
|
||||
]
|
||||
for _model in models_to_check:
|
||||
if _model is not None and _model in litellm.model_cost:
|
||||
cost_info = litellm.model_cost[_model]
|
||||
break
|
||||
if model_info is not None:
|
||||
cost_info = dict(model_info)
|
||||
else:
|
||||
# Build model names for cost lookup
|
||||
base_model_name = model
|
||||
model_name_without_custom_llm_provider: Optional[str] = None
|
||||
if custom_llm_provider and model.startswith(f"{custom_llm_provider}/"):
|
||||
model_name_without_custom_llm_provider = model.replace(
|
||||
f"{custom_llm_provider}/", ""
|
||||
)
|
||||
base_model_name = (
|
||||
f"{custom_llm_provider}/{model_name_without_custom_llm_provider}"
|
||||
)
|
||||
|
||||
verbose_logger.debug(f"Looking up cost for video model: {base_model_name}")
|
||||
|
||||
model_without_provider = model.split("/")[-1]
|
||||
|
||||
# Try model with provider first, fall back to base model name
|
||||
models_to_check: List[Optional[str]] = [
|
||||
base_model_name,
|
||||
model,
|
||||
model_without_provider,
|
||||
model_name_without_custom_llm_provider,
|
||||
]
|
||||
for _model in models_to_check:
|
||||
if _model is not None and _model in litellm.model_cost:
|
||||
cost_info = litellm.model_cost[_model]
|
||||
break
|
||||
|
||||
# If still not found, try with custom_llm_provider prefix
|
||||
if cost_info is None and custom_llm_provider:
|
||||
prefixed_model = f"{custom_llm_provider}/{model}"
|
||||
if prefixed_model in litellm.model_cost:
|
||||
cost_info = litellm.model_cost[prefixed_model]
|
||||
|
||||
# If still not found, try with custom_llm_provider prefix
|
||||
if cost_info is None and custom_llm_provider:
|
||||
prefixed_model = f"{custom_llm_provider}/{model}"
|
||||
if prefixed_model in litellm.model_cost:
|
||||
cost_info = litellm.model_cost[prefixed_model]
|
||||
if cost_info is None:
|
||||
raise Exception(
|
||||
f"Model not found in cost map. Tried checking {models_to_check}"
|
||||
f"Model not found in cost map for model={model}"
|
||||
)
|
||||
|
||||
# Check for video-specific cost per second first
|
||||
|
|
|
|||
|
|
@ -234,6 +234,8 @@ class BaseAWSLLM:
|
|||
aws_session_token=aws_session_token,
|
||||
aws_role_name=aws_role_name,
|
||||
aws_session_name=aws_session_name,
|
||||
aws_region_name=aws_region_name,
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
aws_external_id=aws_external_id,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
|
@ -733,6 +735,7 @@ class BaseAWSLLM:
|
|||
region: str,
|
||||
web_identity_token_file: str,
|
||||
aws_external_id: Optional[str] = None,
|
||||
aws_sts_endpoint: Optional[str] = None,
|
||||
ssl_verify: Optional[Union[bool, str]] = None,
|
||||
) -> dict:
|
||||
"""Handle cross-account role assumption for IRSA."""
|
||||
|
|
@ -744,11 +747,13 @@ class BaseAWSLLM:
|
|||
with open(web_identity_token_file, "r") as f:
|
||||
web_identity_token = f.read().strip()
|
||||
|
||||
irsa_sts_kwargs: dict = {"region_name": region, "verify": self._get_ssl_verify(ssl_verify)}
|
||||
if aws_sts_endpoint is not None:
|
||||
irsa_sts_kwargs["endpoint_url"] = aws_sts_endpoint
|
||||
|
||||
# Create an STS client without credentials
|
||||
with tracer.trace("boto3.client(sts) for manual IRSA"):
|
||||
sts_client = boto3.client(
|
||||
"sts", region_name=region, verify=self._get_ssl_verify(ssl_verify)
|
||||
)
|
||||
sts_client = boto3.client("sts", **irsa_sts_kwargs)
|
||||
|
||||
# Manually assume the IRSA role with the session name
|
||||
verbose_logger.debug(
|
||||
|
|
@ -767,11 +772,10 @@ class BaseAWSLLM:
|
|||
with tracer.trace("boto3.client(sts) with manual IRSA credentials"):
|
||||
sts_client_with_creds = boto3.client(
|
||||
"sts",
|
||||
region_name=region,
|
||||
aws_access_key_id=irsa_creds["AccessKeyId"],
|
||||
aws_secret_access_key=irsa_creds["SecretAccessKey"],
|
||||
aws_session_token=irsa_creds["SessionToken"],
|
||||
verify=self._get_ssl_verify(ssl_verify),
|
||||
**irsa_sts_kwargs,
|
||||
)
|
||||
|
||||
# Get current caller identity for debugging
|
||||
|
|
@ -804,16 +808,19 @@ class BaseAWSLLM:
|
|||
aws_session_name: str,
|
||||
region: str,
|
||||
aws_external_id: Optional[str] = None,
|
||||
aws_sts_endpoint: Optional[str] = None,
|
||||
ssl_verify: Optional[Union[bool, str]] = None,
|
||||
) -> dict:
|
||||
"""Handle same-account role assumption for IRSA."""
|
||||
import boto3
|
||||
|
||||
irsa_sts_kwargs: dict = {"region_name": region, "verify": self._get_ssl_verify(ssl_verify)}
|
||||
if aws_sts_endpoint is not None:
|
||||
irsa_sts_kwargs["endpoint_url"] = aws_sts_endpoint
|
||||
|
||||
verbose_logger.debug("Same account role assumption, using automatic IRSA")
|
||||
with tracer.trace("boto3.client(sts) with automatic IRSA"):
|
||||
sts_client = boto3.client(
|
||||
"sts", region_name=region, verify=self._get_ssl_verify(ssl_verify)
|
||||
)
|
||||
sts_client = boto3.client("sts", **irsa_sts_kwargs)
|
||||
|
||||
# Get current caller identity for debugging
|
||||
try:
|
||||
|
|
@ -867,6 +874,8 @@ class BaseAWSLLM:
|
|||
aws_session_token: Optional[str],
|
||||
aws_role_name: str,
|
||||
aws_session_name: str,
|
||||
aws_region_name: Optional[str] = None,
|
||||
aws_sts_endpoint: Optional[str] = None,
|
||||
aws_external_id: Optional[str] = None,
|
||||
ssl_verify: Optional[Union[bool, str]] = None,
|
||||
) -> Tuple[Credentials, Optional[int]]:
|
||||
|
|
@ -880,6 +889,8 @@ class BaseAWSLLM:
|
|||
web_identity_token_file = os.getenv("AWS_WEB_IDENTITY_TOKEN_FILE")
|
||||
irsa_role_arn = os.getenv("AWS_ROLE_ARN")
|
||||
|
||||
region = aws_region_name or os.getenv("AWS_REGION") or os.getenv("AWS_DEFAULT_REGION")
|
||||
|
||||
# If we have IRSA environment variables and no explicit credentials,
|
||||
# we need to use the web identity token flow
|
||||
if (
|
||||
|
|
@ -895,12 +906,8 @@ class BaseAWSLLM:
|
|||
)
|
||||
|
||||
try:
|
||||
# Get region from environment
|
||||
region = (
|
||||
os.getenv("AWS_REGION")
|
||||
or os.getenv("AWS_DEFAULT_REGION")
|
||||
or "us-east-1"
|
||||
)
|
||||
# Use passed-in region when set, else env, else default (align with AssumeRole path)
|
||||
region = region or "us-east-1"
|
||||
|
||||
# Check if we need to do cross-account role assumption
|
||||
if aws_role_name != irsa_role_arn:
|
||||
|
|
@ -911,6 +918,7 @@ class BaseAWSLLM:
|
|||
region,
|
||||
web_identity_token_file,
|
||||
aws_external_id,
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
else:
|
||||
|
|
@ -919,6 +927,7 @@ class BaseAWSLLM:
|
|||
aws_session_name,
|
||||
region,
|
||||
aws_external_id,
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
|
|
@ -940,11 +949,14 @@ class BaseAWSLLM:
|
|||
|
||||
# In EKS/IRSA environments, use ambient credentials (no explicit keys needed)
|
||||
# This allows the web identity token to work automatically
|
||||
sts_client_kwargs: dict = {"verify": self._get_ssl_verify(ssl_verify)}
|
||||
if region is not None:
|
||||
sts_client_kwargs["region_name"] = region
|
||||
if aws_sts_endpoint is not None:
|
||||
sts_client_kwargs["endpoint_url"] = aws_sts_endpoint
|
||||
if aws_access_key_id is None and aws_secret_access_key is None:
|
||||
with tracer.trace("boto3.client(sts)"):
|
||||
sts_client = boto3.client(
|
||||
"sts", verify=self._get_ssl_verify(ssl_verify)
|
||||
)
|
||||
sts_client = boto3.client("sts", **sts_client_kwargs)
|
||||
else:
|
||||
with tracer.trace("boto3.client(sts)"):
|
||||
sts_client = boto3.client(
|
||||
|
|
@ -952,7 +964,7 @@ class BaseAWSLLM:
|
|||
aws_access_key_id=aws_access_key_id,
|
||||
aws_secret_access_key=aws_secret_access_key,
|
||||
aws_session_token=aws_session_token,
|
||||
verify=self._get_ssl_verify(ssl_verify),
|
||||
**sts_client_kwargs,
|
||||
)
|
||||
|
||||
assume_role_params = {
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ import httpx
|
|||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
from litellm.llms.bedrock.common_utils import BedrockError
|
||||
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
from litellm.passthrough.utils import CommonUtils
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -94,6 +95,9 @@ class AmazonBedrockOpenAIConfig(OpenAIGPTConfig, BaseAWSLLM):
|
|||
aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint,
|
||||
aws_region_name=aws_region_name,
|
||||
)
|
||||
|
||||
# Encode model ID for ARNs (e.g., :imported-model/ -> :imported-model%2F)
|
||||
model_id = CommonUtils.encode_bedrock_runtime_modelid_arn(model_id)
|
||||
|
||||
# Build the invoke URL
|
||||
if stream:
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ from typing import Literal, Optional, Tuple
|
|||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
|
||||
from litellm.types.utils import CallTypes, Usage
|
||||
from litellm.types.utils import CallTypes, ModelInfo, Usage
|
||||
from litellm.utils import get_model_info
|
||||
|
||||
|
||||
|
|
@ -129,7 +129,10 @@ def cost_per_second(
|
|||
|
||||
|
||||
def video_generation_cost(
|
||||
model: str, duration_seconds: float, custom_llm_provider: Optional[str] = None
|
||||
model: str,
|
||||
duration_seconds: float,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
model_info: Optional[ModelInfo] = None,
|
||||
) -> float:
|
||||
"""
|
||||
Calculates the cost for video generation based on duration in seconds.
|
||||
|
|
@ -138,14 +141,18 @@ def video_generation_cost(
|
|||
- model: str, the model name without provider prefix
|
||||
- duration_seconds: float, the duration of the generated video in seconds
|
||||
- custom_llm_provider: str, the custom llm provider
|
||||
- model_info: Optional[dict], deployment-level model info containing
|
||||
custom video pricing. When provided, skips the global
|
||||
get_model_info() lookup so that deployment-specific pricing is used.
|
||||
|
||||
Returns:
|
||||
float - total_cost_in_usd
|
||||
"""
|
||||
## GET MODEL INFO
|
||||
model_info = get_model_info(
|
||||
model=model, custom_llm_provider=custom_llm_provider or "openai"
|
||||
)
|
||||
if model_info is None:
|
||||
model_info = get_model_info(
|
||||
model=model, custom_llm_provider=custom_llm_provider or "openai"
|
||||
)
|
||||
|
||||
# Check for video-specific cost per second
|
||||
video_cost_per_second = model_info.get("output_cost_per_video_per_second")
|
||||
|
|
|
|||
|
|
@ -3517,7 +3517,7 @@ def test_bedrock_openai_imported_model():
|
|||
print(f"URL: {url}")
|
||||
assert "bedrock-runtime.us-east-1.amazonaws.com" in url
|
||||
assert (
|
||||
"arn:aws:bedrock:us-east-1:117159858402:imported-model/m4gc1mrfuddy" in url
|
||||
"arn:aws:bedrock:us-east-1:117159858402:imported-model%2Fm4gc1mrfuddy" in url
|
||||
)
|
||||
assert "/invoke" in url
|
||||
|
||||
|
|
|
|||
|
|
@ -541,25 +541,35 @@ def test_different_roles_without_session_names_should_not_share_cache():
|
|||
assert cache_key1 != cache_key2
|
||||
|
||||
|
||||
def test_eks_irsa_ambient_credentials_used():
|
||||
@pytest.mark.parametrize(
|
||||
"role_kwargs,expected_client_kwargs",
|
||||
[
|
||||
({}, {"verify": True}),
|
||||
({"aws_region_name": "us-east-1"}, {"region_name": "us-east-1", "verify": True}),
|
||||
(
|
||||
{"aws_sts_endpoint": "https://sts.eu-west-1.amazonaws.com"},
|
||||
{"endpoint_url": "https://sts.eu-west-1.amazonaws.com", "verify": True},
|
||||
),
|
||||
],
|
||||
ids=["no_region_or_endpoint", "regional_sts", "explicit_sts_endpoint"],
|
||||
)
|
||||
def test_eks_irsa_ambient_credentials_used(role_kwargs, expected_client_kwargs):
|
||||
"""
|
||||
Test that in EKS/IRSA environments, ambient credentials are used when no explicit keys provided.
|
||||
This allows web identity tokens to work automatically.
|
||||
"""
|
||||
# Isolate from ambient AWS_REGION/AWS_DEFAULT_REGION so no_region_or_endpoint is deterministic
|
||||
env_without_aws_region = {
|
||||
k: v
|
||||
for k, v in os.environ.items()
|
||||
if k not in ("AWS_REGION", "AWS_DEFAULT_REGION")
|
||||
}
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
|
||||
# Mock the boto3 STS client
|
||||
mock_sts_client = MagicMock()
|
||||
|
||||
# Mock the STS response with proper expiration handling
|
||||
mock_expiry = MagicMock()
|
||||
mock_expiry.tzinfo = timezone.utc
|
||||
current_time = datetime.now(timezone.utc)
|
||||
# Create a timedelta object that returns 3600 when total_seconds() is called
|
||||
time_diff = MagicMock()
|
||||
time_diff.total_seconds.return_value = 3600
|
||||
mock_expiry.__sub__ = MagicMock(return_value=time_diff)
|
||||
|
||||
mock_sts_response = {
|
||||
"Credentials": {
|
||||
"AccessKeyId": "assumed-access-key",
|
||||
|
|
@ -568,54 +578,82 @@ def test_eks_irsa_ambient_credentials_used():
|
|||
"Expiration": mock_expiry,
|
||||
}
|
||||
}
|
||||
mock_sts_client = MagicMock()
|
||||
mock_sts_client.assume_role.return_value = mock_sts_response
|
||||
|
||||
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
|
||||
|
||||
# Call with no explicit credentials (EKS/IRSA scenario)
|
||||
credentials, ttl = base_aws_llm._auth_with_aws_role(
|
||||
aws_access_key_id=None,
|
||||
aws_secret_access_key=None,
|
||||
aws_session_token=None,
|
||||
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
aws_session_name="test-session"
|
||||
)
|
||||
|
||||
# Should create STS client without explicit credentials (using ambient credentials)
|
||||
# Note: verify parameter is passed for SSL verification
|
||||
mock_boto3_client.assert_called_once_with("sts", verify=True)
|
||||
|
||||
# Should call assume_role
|
||||
mock_sts_client.assume_role.assert_called_once_with(
|
||||
RoleArn="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
RoleSessionName="test-session"
|
||||
)
|
||||
|
||||
# Verify credentials are returned correctly
|
||||
assert credentials.access_key == "assumed-access-key"
|
||||
assert credentials.secret_key == "assumed-secret-key"
|
||||
assert credentials.token == "assumed-session-token"
|
||||
assert ttl is not None
|
||||
|
||||
with patch.dict(os.environ, env_without_aws_region, clear=True):
|
||||
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
|
||||
credentials, ttl = base_aws_llm._auth_with_aws_role(
|
||||
aws_access_key_id=None,
|
||||
aws_secret_access_key=None,
|
||||
aws_session_token=None,
|
||||
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
aws_session_name="test-session",
|
||||
**role_kwargs,
|
||||
)
|
||||
mock_boto3_client.assert_called_once_with(
|
||||
"sts", **expected_client_kwargs
|
||||
)
|
||||
mock_sts_client.assume_role.assert_called_once_with(
|
||||
RoleArn="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
RoleSessionName="test-session",
|
||||
)
|
||||
assert credentials.access_key == "assumed-access-key"
|
||||
assert ttl is not None
|
||||
|
||||
|
||||
def test_explicit_credentials_used_when_provided():
|
||||
@pytest.mark.parametrize(
|
||||
"role_kwargs,expected_client_kwargs",
|
||||
[
|
||||
(
|
||||
{},
|
||||
{
|
||||
"aws_access_key_id": "explicit-access-key",
|
||||
"aws_secret_access_key": "explicit-secret-key",
|
||||
"aws_session_token": "assumed-session-token",
|
||||
"verify": True,
|
||||
},
|
||||
),
|
||||
(
|
||||
{"aws_region_name": "us-east-1"},
|
||||
{
|
||||
"region_name": "us-east-1",
|
||||
"aws_access_key_id": "explicit-access-key",
|
||||
"aws_secret_access_key": "explicit-secret-key",
|
||||
"aws_session_token": "assumed-session-token",
|
||||
"verify": True,
|
||||
},
|
||||
),
|
||||
(
|
||||
{"aws_sts_endpoint": "https://sts.eu-west-1.amazonaws.com"},
|
||||
{
|
||||
"endpoint_url": "https://sts.eu-west-1.amazonaws.com",
|
||||
"aws_access_key_id": "explicit-access-key",
|
||||
"aws_secret_access_key": "explicit-secret-key",
|
||||
"aws_session_token": "assumed-session-token",
|
||||
"verify": True,
|
||||
},
|
||||
),
|
||||
],
|
||||
ids=["no_region_or_endpoint", "regional_sts", "explicit_sts_endpoint"],
|
||||
)
|
||||
def test_explicit_credentials_used_when_provided(role_kwargs, expected_client_kwargs):
|
||||
"""
|
||||
Test that explicit credentials are used when provided (non-EKS/IRSA scenario).
|
||||
"""
|
||||
# Isolate from ambient AWS_REGION/AWS_DEFAULT_REGION so no_region_or_endpoint is deterministic
|
||||
env_without_aws_region = {
|
||||
k: v
|
||||
for k, v in os.environ.items()
|
||||
if k not in ("AWS_REGION", "AWS_DEFAULT_REGION")
|
||||
}
|
||||
base_aws_llm = BaseAWSLLM()
|
||||
|
||||
# Mock the boto3 STS client
|
||||
mock_sts_client = MagicMock()
|
||||
|
||||
# Mock the STS response with proper expiration handling
|
||||
mock_expiry = MagicMock()
|
||||
mock_expiry.tzinfo = timezone.utc
|
||||
current_time = datetime.now(timezone.utc)
|
||||
# Create a timedelta object that returns 3600 when total_seconds() is called
|
||||
time_diff = MagicMock()
|
||||
time_diff.total_seconds.return_value = 3600
|
||||
mock_expiry.__sub__ = MagicMock(return_value=time_diff)
|
||||
|
||||
mock_sts_response = {
|
||||
"Credentials": {
|
||||
"AccessKeyId": "assumed-access-key",
|
||||
|
|
@ -624,40 +662,30 @@ def test_explicit_credentials_used_when_provided():
|
|||
"Expiration": mock_expiry,
|
||||
}
|
||||
}
|
||||
mock_sts_client = MagicMock()
|
||||
mock_sts_client.assume_role.return_value = mock_sts_response
|
||||
|
||||
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
|
||||
|
||||
# Call with explicit credentials
|
||||
credentials, ttl = base_aws_llm._auth_with_aws_role(
|
||||
aws_access_key_id="explicit-access-key",
|
||||
aws_secret_access_key="explicit-secret-key",
|
||||
aws_session_token="assumed-session-token",
|
||||
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
aws_session_name="test-session"
|
||||
)
|
||||
|
||||
# Should create STS client with explicit credentials
|
||||
# Note: verify parameter is passed for SSL verification
|
||||
mock_boto3_client.assert_called_once_with(
|
||||
"sts",
|
||||
aws_access_key_id="explicit-access-key",
|
||||
aws_secret_access_key="explicit-secret-key",
|
||||
aws_session_token="assumed-session-token",
|
||||
verify=True,
|
||||
)
|
||||
|
||||
# Should call assume_role
|
||||
mock_sts_client.assume_role.assert_called_once_with(
|
||||
RoleArn="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
RoleSessionName="test-session"
|
||||
)
|
||||
|
||||
# Verify credentials are returned correctly
|
||||
assert credentials.access_key == "assumed-access-key"
|
||||
assert credentials.secret_key == "assumed-secret-key"
|
||||
assert credentials.token == "assumed-session-token"
|
||||
assert ttl is not None
|
||||
|
||||
with patch.dict(os.environ, env_without_aws_region, clear=True):
|
||||
with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client:
|
||||
credentials, ttl = base_aws_llm._auth_with_aws_role(
|
||||
aws_access_key_id="explicit-access-key",
|
||||
aws_secret_access_key="explicit-secret-key",
|
||||
aws_session_token="assumed-session-token",
|
||||
aws_role_name="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
aws_session_name="test-session",
|
||||
**role_kwargs,
|
||||
)
|
||||
mock_boto3_client.assert_called_once_with(
|
||||
"sts", **expected_client_kwargs
|
||||
)
|
||||
mock_sts_client.assume_role.assert_called_once_with(
|
||||
RoleArn="arn:aws:iam::2222222222222:role/LitellmEvalBedrockRole",
|
||||
RoleSessionName="test-session",
|
||||
)
|
||||
assert credentials.access_key == "assumed-access-key"
|
||||
assert credentials.secret_key == "assumed-secret-key"
|
||||
assert credentials.token == "assumed-session-token"
|
||||
assert ttl is not None
|
||||
|
||||
|
||||
def test_partial_credentials_still_use_ambient():
|
||||
|
|
|
|||
|
|
@ -242,6 +242,77 @@ class TestVideoGeneration:
|
|||
custom_llm_provider="openai"
|
||||
)
|
||||
|
||||
def test_video_generation_cost_with_custom_model_info(self):
|
||||
"""Test that custom model_info pricing is applied for video generation.
|
||||
|
||||
When a deployment has custom pricing via model_info, it should be used
|
||||
instead of looking up the global litellm.model_cost map.
|
||||
|
||||
Related: https://github.com/BerriAI/litellm/issues/21907
|
||||
"""
|
||||
model_info = {
|
||||
"output_cost_per_video_per_second": 0.05,
|
||||
}
|
||||
cost = default_video_cost_calculator(
|
||||
model="my-custom-video-model",
|
||||
duration_seconds=10.0,
|
||||
model_info=model_info,
|
||||
)
|
||||
assert cost == 0.5
|
||||
|
||||
def test_video_generation_cost_custom_model_info_fallback_to_per_second(self):
|
||||
"""Test that output_cost_per_second is used as fallback when
|
||||
output_cost_per_video_per_second is not set in custom model_info.
|
||||
|
||||
Related: https://github.com/BerriAI/litellm/issues/21907
|
||||
"""
|
||||
model_info = {
|
||||
"output_cost_per_second": 0.10,
|
||||
}
|
||||
cost = default_video_cost_calculator(
|
||||
model="my-custom-video-model",
|
||||
duration_seconds=5.0,
|
||||
model_info=model_info,
|
||||
)
|
||||
assert cost == 0.5
|
||||
|
||||
def test_video_generation_cost_custom_pricing_through_completion_cost(self):
|
||||
"""Test that custom video pricing flows through completion_cost via litellm_logging_obj.
|
||||
|
||||
This tests the full cost calculation path: completion_cost extracts model_info
|
||||
from litellm_logging_obj.litellm_params.metadata.model_info and passes it to
|
||||
the video cost calculator.
|
||||
|
||||
Related: https://github.com/BerriAI/litellm/issues/21907
|
||||
"""
|
||||
from litellm.cost_calculator import completion_cost
|
||||
|
||||
# Create mock response with usage containing duration_seconds
|
||||
mock_response = MagicMock()
|
||||
mock_response.usage = MagicMock()
|
||||
mock_response.usage.duration_seconds = 10.0
|
||||
type(mock_response)._hidden_params = {}
|
||||
|
||||
# Create mock litellm_logging_obj with custom pricing
|
||||
mock_logging_obj = MagicMock()
|
||||
mock_logging_obj.litellm_params = {
|
||||
"metadata": {
|
||||
"model_info": {
|
||||
"output_cost_per_video_per_second": 0.05,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
cost = completion_cost(
|
||||
completion_response=mock_response,
|
||||
model="openai/hunyuanvideo",
|
||||
call_type="create_video",
|
||||
custom_llm_provider="openai",
|
||||
custom_pricing=True,
|
||||
litellm_logging_obj=mock_logging_obj,
|
||||
)
|
||||
assert cost == 0.5
|
||||
|
||||
def test_video_generation_with_files(self):
|
||||
"""Test video generation with file uploads."""
|
||||
config = OpenAIVideoConfig()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue