Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_batch_output_single_pass

# Conflicts:
#	basedpyright-code-budget.json
#	ruff-strict-budget.json
#	type-discipline-budget.json
This commit is contained in:
mateo-berri 2026-07-30 10:53:20 -07:00
commit a97233067d
37 changed files with 1164 additions and 1336 deletions

View file

@ -1,12 +1,12 @@
{
"reportAny": {
"limit": 33210
"limit": 31903
},
"reportArgumentType": {
"limit": 2648
"limit": 2645
},
"reportAssignmentType": {
"limit": 330
"limit": 329
},
"reportAttributeAccessIssue": {
"limit": 516
@ -18,13 +18,13 @@
"limit": 59
},
"reportDeprecated": {
"limit": 326
"limit": 325
},
"reportDuplicateImport": {
"limit": 42
},
"reportExplicitAny": {
"limit": 10228
"limit": 10214
},
"reportFunctionMemberAccess": {
"limit": 11
@ -54,10 +54,10 @@
"limit": 0
},
"reportMissingParameterType": {
"limit": 5893
"limit": 5869
},
"reportMissingTypeArgument": {
"limit": 15883
"limit": 15861
},
"reportMissingTypeStubs": {
"limit": 41
@ -72,7 +72,7 @@
"limit": 0
},
"reportOptionalMemberAccess": {
"limit": 1085
"limit": 1079
},
"reportOptionalOperand": {
"limit": 0
@ -84,13 +84,13 @@
"limit": 77
},
"reportPrivateUsage": {
"limit": 2438
"limit": 2437
},
"reportRedeclaration": {
"limit": 12
},
"reportReturnType": {
"limit": 225
"limit": 219
},
"reportTypedDictNotRequiredAccess": {
"limit": 27
@ -99,31 +99,31 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 45567
"limit": 45366
},
"reportUnknownLambdaType": {
"limit": 113
},
"reportUnknownMemberType": {
"limit": 40523
"limit": 40477
},
"reportUnknownParameterType": {
"limit": 20381
"limit": 20338
},
"reportUnknownVariableType": {
"limit": 32095
"limit": 32047
},
"reportUnnecessaryCast": {
"limit": 177
},
"reportUnnecessaryComparison": {
"limit": 1022
"limit": 1021
},
"reportUnnecessaryContains": {
"limit": 7
},
"reportUnnecessaryIsInstance": {
"limit": 1206
"limit": 1205
},
"reportUntypedBaseClass": {
"limit": 165
@ -135,10 +135,10 @@
"limit": 33
},
"reportUnusedFunction": {
"limit": 205
"limit": 204
},
"reportUnusedImport": {
"limit": 1005
"limit": 1003
},
"reportUnusedVariable": {
"limit": 1297

View file

@ -0,0 +1,523 @@
{
"annotations": {
"list": []
},
"editable": true,
"fiscalYearStartMonth": 0,
"graphTooltip": 0,
"links": [],
"panels": [
{
"type": "stat",
"title": "Requests",
"gridPos": {
"h": 4,
"w": 6,
"x": 0,
"y": 0
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "short",
"decimals": 0,
"color": {
"mode": "fixed",
"fixedColor": "blue"
}
},
"overrides": []
},
"options": {
"reduceOptions": {
"calcs": [
"lastNotNull"
],
"fields": "",
"values": false
},
"colorMode": "background",
"graphMode": "none"
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"instant": true,
"expr": "sum(increase(gen_ai_client_operation_duration_seconds_count{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range]))"
}
],
"id": 1
},
{
"type": "stat",
"title": "Spend",
"description": "LiteLLM's computed cost for the selected window, from gen_ai.usage.cost",
"gridPos": {
"h": 4,
"w": 6,
"x": 6,
"y": 0
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "currencyUSD",
"decimals": 4,
"color": {
"mode": "fixed",
"fixedColor": "green"
}
},
"overrides": []
},
"options": {
"reduceOptions": {
"calcs": [
"lastNotNull"
],
"fields": "",
"values": false
},
"colorMode": "background",
"graphMode": "none"
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"instant": true,
"expr": "sum(increase(gen_ai_usage_cost_USD_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range]))"
}
],
"id": 2
},
{
"type": "stat",
"title": "Tokens",
"gridPos": {
"h": 4,
"w": 6,
"x": 12,
"y": 0
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "short",
"decimals": 0,
"color": {
"mode": "fixed",
"fixedColor": "purple"
}
},
"overrides": []
},
"options": {
"reduceOptions": {
"calcs": [
"lastNotNull"
],
"fields": "",
"values": false
},
"colorMode": "background",
"graphMode": "none"
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"instant": true,
"expr": "sum(increase(gen_ai_client_token_usage_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range]))"
}
],
"id": 3
},
{
"type": "stat",
"title": "p95 request duration",
"gridPos": {
"h": 4,
"w": 6,
"x": 18,
"y": 0
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "s",
"decimals": 2,
"color": {
"mode": "fixed",
"fixedColor": "orange"
}
},
"overrides": []
},
"options": {
"reduceOptions": {
"calcs": [
"lastNotNull"
],
"fields": "",
"values": false
},
"colorMode": "background",
"graphMode": "none"
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"instant": true,
"expr": "histogram_quantile(0.95, sum by (le) (rate(gen_ai_client_operation_duration_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range])))"
}
],
"id": 4
},
{
"type": "timeseries",
"title": "Request rate by model",
"gridPos": {
"h": 8,
"w": 12,
"x": 0,
"y": 4
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "reqpm",
"custom": {
"lineWidth": 2,
"fillOpacity": 8,
"showPoints": "never"
}
},
"overrides": []
},
"options": {
"legend": {
"displayMode": "list",
"placement": "bottom"
},
"tooltip": {
"mode": "multi",
"sort": "desc"
}
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"legendFormat": "{{gen_ai_request_model}}",
"expr": "sum by (gen_ai_request_model) (rate(gen_ai_client_operation_duration_seconds_count{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])) * 60"
}
],
"id": 5
},
{
"type": "timeseries",
"title": "Spend rate by model",
"description": "USD per hour, derived from the gen_ai.usage.cost histogram",
"gridPos": {
"h": 8,
"w": 12,
"x": 12,
"y": 4
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "currencyUSD",
"custom": {
"lineWidth": 2,
"fillOpacity": 8,
"showPoints": "never"
}
},
"overrides": []
},
"options": {
"legend": {
"displayMode": "list",
"placement": "bottom"
},
"tooltip": {
"mode": "multi",
"sort": "desc"
}
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"legendFormat": "{{gen_ai_request_model}}",
"expr": "sum by (gen_ai_request_model) (rate(gen_ai_usage_cost_USD_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])) * 3600"
}
],
"id": 6
},
{
"type": "timeseries",
"title": "Tokens per minute by model and type",
"description": "gen_ai.client.token.usage split by the gen_ai.token.type attribute",
"gridPos": {
"h": 8,
"w": 12,
"x": 0,
"y": 12
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "short",
"custom": {
"lineWidth": 2,
"fillOpacity": 8,
"showPoints": "never"
}
},
"overrides": []
},
"options": {
"legend": {
"displayMode": "list",
"placement": "bottom"
},
"tooltip": {
"mode": "multi",
"sort": "desc"
}
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"legendFormat": "{{gen_ai_request_model}} {{gen_ai_token_type}}",
"expr": "sum by (gen_ai_request_model, gen_ai_token_type) (rate(gen_ai_client_token_usage_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])) * 60"
}
],
"id": 7
},
{
"type": "timeseries",
"title": "p95 request duration by model",
"gridPos": {
"h": 8,
"w": 12,
"x": 12,
"y": 12
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "s",
"custom": {
"lineWidth": 2,
"fillOpacity": 0,
"showPoints": "never"
}
},
"overrides": []
},
"options": {
"legend": {
"displayMode": "list",
"placement": "bottom"
},
"tooltip": {
"mode": "multi",
"sort": "desc"
}
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"legendFormat": "{{gen_ai_request_model}}",
"expr": "histogram_quantile(0.95, sum by (le, gen_ai_request_model) (rate(gen_ai_client_operation_duration_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])))"
}
],
"id": 8
},
{
"type": "timeseries",
"title": "p95 time to first token (streaming)",
"description": "gen_ai.server.time_to_first_token, recorded only for streaming requests",
"gridPos": {
"h": 8,
"w": 12,
"x": 0,
"y": 20
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "s",
"custom": {
"lineWidth": 2,
"fillOpacity": 0,
"showPoints": "never"
}
},
"overrides": []
},
"options": {
"legend": {
"displayMode": "list",
"placement": "bottom"
},
"tooltip": {
"mode": "multi",
"sort": "desc"
}
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"legendFormat": "{{gen_ai_request_model}}",
"expr": "histogram_quantile(0.95, sum by (le, gen_ai_request_model) (rate(gen_ai_server_time_to_first_token_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])))"
}
],
"id": 9
},
{
"type": "timeseries",
"title": "p95 provider generation time",
"description": "gen_ai.client.response.duration, upstream generation time excluding LiteLLM overhead",
"gridPos": {
"h": 8,
"w": 12,
"x": 12,
"y": 20
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "s",
"custom": {
"lineWidth": 2,
"fillOpacity": 0,
"showPoints": "never"
}
},
"overrides": []
},
"options": {
"legend": {
"displayMode": "list",
"placement": "bottom"
},
"tooltip": {
"mode": "multi",
"sort": "desc"
}
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"legendFormat": "{{gen_ai_request_model}}",
"expr": "histogram_quantile(0.95, sum by (le, gen_ai_request_model) (rate(gen_ai_client_response_duration_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])))"
}
],
"id": 10
}
],
"preload": false,
"refresh": "30s",
"schemaVersion": 42,
"tags": [
"litellm",
"genai",
"opentelemetry"
],
"templating": {
"list": [
{
"name": "datasource",
"label": "Prometheus",
"type": "datasource",
"query": "prometheus",
"current": {},
"hide": 0
},
{
"name": "service",
"label": "Service",
"type": "query",
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"query": "label_values(gen_ai_client_operation_duration_seconds_count, service_name)",
"refresh": 2,
"includeAll": true,
"multi": true,
"current": {
"text": "All",
"value": "$__all"
}
},
{
"name": "model",
"label": "Model",
"type": "query",
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"query": "label_values(gen_ai_client_operation_duration_seconds_count{service_name=~\"$service\"}, gen_ai_request_model)",
"refresh": 2,
"includeAll": true,
"multi": true,
"current": {
"text": "All",
"value": "$__all"
}
}
]
},
"time": {
"from": "now-1h",
"to": "now"
},
"timepicker": {},
"timezone": "browser",
"title": "LiteLLM GenAI (OpenTelemetry)",
"uid": "litellm-genai-otel",
"weekStart": ""
}

View file

@ -0,0 +1,35 @@
# LiteLLM GenAI dashboard (OpenTelemetry metrics)
Dashboard for the `gen_ai.*` metrics the OpenTelemetry v2 integration emits, as opposed to the `litellm_*` Prometheus metrics the other dashboards in this folder chart.
Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source. Panels: request count, spend, token count, p95 duration, request rate by model, spend rate per hour by model, tokens per minute split by input and output, p95 duration by model, p95 time to first token, and p95 provider generation time. Template variables for data source, service, and model.
## Pre-requisites
Metrics are off by default. In the proxy environment:
```shell
LITELLM_OTEL_V2=true
LITELLM_OTEL_INTEGRATION_ENABLE_METRICS=true
OTEL_EXPORTER="otlp_http"
OTEL_ENDPOINT="<your OTLP endpoint>"
```
You also need the metric attribute filter, or the panels will plot flat lines at zero. LiteLLM's default attribute set includes per-request fields, so nearly every request lands in its own time series with a single sample, and `rate()` has nothing to compute over:
```yaml title="config.yaml"
callback_settings:
otel:
attributes:
include_list:
- gen_ai.operation.name
- gen_ai.system
- gen_ai.request.model
- gen_ai.framework
```
See [Grafana Cloud](https://docs.litellm.ai/docs/observability/grafana_cloud) for the full setup, and [OpenTelemetry v2](https://docs.litellm.ai/docs/observability/opentelemetry_v2#metrics) for the metric reference.
## Note on Grafana's AI Observability integration
Grafana Cloud ships prebuilt GenAI dashboards that query these same metric names, so they look like a drop-in alternative to this one. They are not: twenty of their twenty-two panels filter on `telemetry_sdk_name="openlit"`, a label LiteLLM does not carry and cannot be configured to add, so those panels stay empty.

View file

@ -2,6 +2,10 @@
This folder contains the `json` for creating Grafana Dashboards
## [LiteLLM GenAI Dashboard (OpenTelemetry)](./dashboard_genai_otel)
Charts the `gen_ai.*` metrics from the OpenTelemetry v2 integration: spend, tokens, request rate, and latency percentiles by model. Separate from the dashboards below, which chart the `litellm_*` Prometheus metrics.
## [LiteLLM v2 Dashboard](./dashboard_v2)
<img width="1316" alt="grafana_1" src="https://github.com/user-attachments/assets/d0df802d-0cb9-4906-a679-941c547789ab">

View file

@ -5,7 +5,6 @@ from .invoke_handler import (
AmazonAnthropicClaudeStreamDecoder,
AmazonDeepSeekR1StreamDecoder,
AWSEventStreamDecoder,
BedrockLLM,
)

View file

@ -1,19 +1,10 @@
"""
TODO: DELETE FILE. Bedrock LLM is no longer used. Goto `litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py`
"""
import copy
import time
import types
from functools import partial
from typing import (
AsyncIterator,
Callable,
Iterator,
Optional,
Tuple,
cast,
get_args,
)
import httpx # type: ignore
@ -25,16 +16,6 @@ from litellm.caching.caching import InMemoryCache
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
from litellm.litellm_core_utils.core_helpers import map_finish_reason
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.logging_utils import track_llm_api_timing
from litellm.litellm_core_utils.prompt_templates.factory import (
cohere_message_pt,
construct_tool_use_system_prompt,
contains_tag,
custom_prompt,
extract_between_tags,
parse_xml_params,
prompt_factory,
)
from litellm.llms.anthropic.chat.handler import (
ModelResponseIterator as AnthropicModelResponseIterator,
)
@ -64,12 +45,9 @@ from litellm.types.utils import (
StreamingChoices,
Usage,
)
from litellm.utils import CustomStreamWrapper, get_secret
from ..base_aws_llm import BaseAWSLLM
from ..common_utils import (
BedrockError,
ModelResponseIterator,
build_bedrock_stream_error,
get_bedrock_response_stream_shape,
get_bedrock_tool_name,
@ -77,9 +55,6 @@ from ..common_utils import (
bedrock_tool_name_mappings: InMemoryCache = InMemoryCache(max_size_in_memory=50, default_ttl=600)
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import (
AmazonBedrockOpenAIConfig,
)
converse_config = AmazonConverseConfig()
@ -351,932 +326,6 @@ def make_sync_call(
raise BedrockError(status_code=500, message=str(e))
class BedrockLLM(BaseAWSLLM):
"""
Example call
```
curl --location --request POST 'https://bedrock-runtime.{aws_region_name}.amazonaws.com/model/{bedrock_model_name}/invoke' \
--header 'Content-Type: application/json' \
--header 'Accept: application/json' \
--user "$AWS_ACCESS_KEY_ID":"$AWS_SECRET_ACCESS_KEY" \
--aws-sigv4 "aws:amz:us-east-1:bedrock" \
--data-raw '{
"prompt": "Hi",
"temperature": 0,
"p": 0.9,
"max_tokens": 4096
}'
```
"""
def __init__(self) -> None:
super().__init__()
@staticmethod
def is_claude_messages_api_model(model: str) -> bool:
"""
Check if the model uses the Claude Messages API (Claude 3+).
Handles:
- Regional prefixes: eu.anthropic.claude-*, us.anthropic.claude-*
- Claude 3 models: claude-3-haiku, claude-3-sonnet, claude-3-opus, claude-3-5-*, claude-3-7-*
- Claude 4 models: claude-opus-4, claude-sonnet-4, claude-haiku-4
"""
# Normalize model string to lowercase for matching
model_lower = model.lower()
# Claude 3+ indicators (all use Messages API)
messages_api_indicators = [
"claude-3", # Claude 3.x models
"claude-opus-4", # Claude Opus 4
"claude-sonnet-4", # Claude Sonnet 4
"claude-haiku-4", # Claude Haiku 4
]
return any(indicator in model_lower for indicator in messages_api_indicators)
def convert_messages_to_prompt(self, model, messages, provider, custom_prompt_dict) -> Tuple[str, Optional[list]]:
# handle anthropic prompts and amazon titan prompts
prompt = ""
chat_history: Optional[list] = None
## CUSTOM PROMPT
if model in custom_prompt_dict:
# check if the model has a registered custom prompt
model_prompt_details = custom_prompt_dict[model]
prompt = custom_prompt(
role_dict=model_prompt_details["roles"],
initial_prompt_value=model_prompt_details.get("initial_prompt_value", ""),
final_prompt_value=model_prompt_details.get("final_prompt_value", ""),
messages=messages,
)
return prompt, None
## ELSE
if provider == "anthropic" or provider == "amazon":
prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock")
elif provider == "mistral":
prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock")
elif provider == "meta" or provider == "llama":
prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock")
elif provider == "openai":
# OpenAI uses messages directly, no prompt conversion needed
# Return empty prompt as it won't be used
prompt = ""
elif provider == "cohere":
prompt, chat_history = cohere_message_pt(messages=messages)
else:
prompt = ""
for message in messages:
if "role" in message:
if message["role"] == "user":
prompt += f"{message['content']}"
else:
prompt += f"{message['content']}"
else:
prompt += f"{message['content']}"
return prompt, chat_history # type: ignore
def process_response(
self,
model: str,
response: httpx.Response,
model_response: ModelResponse,
stream: Optional[bool],
logging_obj: Logging,
optional_params: dict,
api_key: str,
data: Union[dict, str],
messages: List,
print_verbose,
encoding,
) -> Union[ModelResponse, CustomStreamWrapper]:
provider = self.get_bedrock_invoke_provider(model)
## LOGGING
logging_obj.post_call(
input=messages,
api_key=api_key,
original_response=response.text,
additional_args={"complete_input_dict": data},
)
print_verbose(f"raw model_response: {response.text}")
## RESPONSE OBJECT
try:
completion_response = response.json()
except Exception:
raise BedrockError(message=response.text, status_code=422)
outputText: Optional[str] = None
try:
if provider == "cohere":
if "text" in completion_response:
outputText = completion_response["text"] # type: ignore
elif "generations" in completion_response:
outputText = completion_response["generations"][0]["text"]
model_response.choices[0].finish_reason = map_finish_reason(
completion_response["generations"][0]["finish_reason"]
)
elif provider == "anthropic":
if self.is_claude_messages_api_model(model):
json_schemas: dict = {}
_is_function_call = False
## Handle Tool Calling
if "tools" in optional_params:
_is_function_call = True
for tool in optional_params["tools"]:
json_schemas[tool["function"]["name"]] = tool["function"].get("parameters", None)
outputText = completion_response.get("content")[0].get("text", None)
if outputText is not None and contains_tag("invoke", outputText): # OUTPUT PARSE FUNCTION CALL
function_name = extract_between_tags("tool_name", outputText)[0]
function_arguments_str = extract_between_tags("invoke", outputText)[0].strip()
function_arguments_str = f"<invoke>{function_arguments_str}</invoke>"
function_arguments = parse_xml_params(
function_arguments_str,
json_schema=json_schemas.get(
function_name, None
), # check if we have a json schema for this function name)
)
_message = litellm.Message(
tool_calls=[
{
"id": f"call_{uuid.uuid4()}",
"type": "function",
"function": {
"name": function_name,
"arguments": json.dumps(function_arguments),
},
}
],
content=None,
)
model_response.choices[0].message = _message # type: ignore
model_response._hidden_params["original_response"] = (
outputText # allow user to access raw anthropic tool calling response
)
if _is_function_call is True and stream is not None and stream is True:
print_verbose("INSIDE BEDROCK STREAMING TOOL CALLING CONDITION BLOCK")
# return an iterator
streaming_model_response = ModelResponseStream()
streaming_model_response.choices[0].finish_reason = getattr(
model_response.choices[0], "finish_reason", "stop"
)
# streaming_model_response.choices = [litellm.utils.StreamingChoices()]
streaming_choice = litellm.utils.StreamingChoices()
streaming_choice.index = model_response.choices[0].index
_tool_calls = []
print_verbose(f"type of model_response.choices[0]: {type(model_response.choices[0])}")
print_verbose(f"type of streaming_choice: {type(streaming_choice)}")
if isinstance(model_response.choices[0], litellm.Choices):
if getattr(
model_response.choices[0].message, "tool_calls", None
) is not None and isinstance(model_response.choices[0].message.tool_calls, list):
for tool_call in model_response.choices[0].message.tool_calls:
_tool_call = {**tool_call.dict(), "index": 0}
_tool_calls.append(_tool_call)
delta_obj = Delta(
content=getattr(model_response.choices[0].message, "content", None),
role=model_response.choices[0].message.role,
tool_calls=_tool_calls,
)
streaming_choice.delta = delta_obj
streaming_model_response.choices = [streaming_choice]
completion_stream = ModelResponseIterator(model_response=streaming_model_response)
print_verbose(
"Returns anthropic CustomStreamWrapper with 'cached_response' streaming object"
)
return litellm.CustomStreamWrapper(
completion_stream=completion_stream,
model=model,
custom_llm_provider="cached_response",
logging_obj=logging_obj,
)
model_response.choices[0].finish_reason = map_finish_reason(
completion_response.get("stop_reason", "")
)
_usage = litellm.Usage(
prompt_tokens=completion_response["usage"]["input_tokens"],
completion_tokens=completion_response["usage"]["output_tokens"],
total_tokens=completion_response["usage"]["input_tokens"]
+ completion_response["usage"]["output_tokens"],
)
setattr(model_response, "usage", _usage)
else:
outputText = completion_response["completion"]
model_response.choices[0].finish_reason = completion_response["stop_reason"]
elif provider == "ai21":
outputText = completion_response.get("completions")[0].get("data").get("text")
elif provider == "meta" or provider == "llama":
outputText = completion_response["generation"]
elif provider == "openai":
# OpenAI imported models use OpenAI Chat Completions format
if "choices" in completion_response and len(completion_response["choices"]) > 0:
choice = completion_response["choices"][0]
if "message" in choice:
outputText = choice["message"].get("content")
elif "text" in choice: # fallback for completion format
outputText = choice["text"]
# Set finish reason
if "finish_reason" in choice:
model_response.choices[0].finish_reason = map_finish_reason(choice["finish_reason"])
# Set usage if available
if "usage" in completion_response:
usage = completion_response["usage"]
_usage = litellm.Usage(
prompt_tokens=usage.get("prompt_tokens", 0),
completion_tokens=usage.get("completion_tokens", 0),
total_tokens=usage.get("total_tokens", 0),
)
setattr(model_response, "usage", _usage)
elif provider == "mistral":
outputText = completion_response["outputs"][0]["text"]
model_response.choices[0].finish_reason = completion_response["outputs"][0]["stop_reason"]
else: # amazon titan
outputText = completion_response.get("results")[0].get("outputText")
except Exception as e:
raise BedrockError(
message="Error processing={}, Received error={}".format(response.text, str(e)),
status_code=422,
)
try:
if (
outputText is not None
and len(outputText) > 0
and hasattr(model_response.choices[0], "message")
and getattr(model_response.choices[0].message, "tool_calls", None) # type: ignore
is None
):
model_response.choices[0].message.content = outputText # type: ignore
elif (
hasattr(model_response.choices[0], "message")
and getattr(model_response.choices[0].message, "tool_calls", None) # type: ignore
is not None
):
pass
else:
raise Exception()
except Exception as e:
raise BedrockError(
message="Error parsing received text={}.\nError-{}".format(outputText, str(e)),
status_code=response.status_code,
)
if stream and provider == "ai21":
streaming_model_response = ModelResponseStream()
streaming_model_response.choices[0].finish_reason = model_response.choices[ # type: ignore
0
].finish_reason
# streaming_model_response.choices = [litellm.utils.StreamingChoices()]
streaming_choice = litellm.utils.StreamingChoices()
streaming_choice.index = model_response.choices[0].index
delta_obj = litellm.utils.Delta(
content=getattr(model_response.choices[0].message, "content", None), # type: ignore
role=model_response.choices[0].message.role, # type: ignore
)
streaming_choice.delta = delta_obj
streaming_model_response.choices = [streaming_choice]
mri = ModelResponseIterator(model_response=streaming_model_response)
return CustomStreamWrapper(
completion_stream=mri,
model=model,
custom_llm_provider="cached_response",
logging_obj=logging_obj,
)
## CALCULATING USAGE - bedrock returns usage in the headers
# Skip if usage was already set (e.g., from JSON response for OpenAI provider)
if not hasattr(model_response, "usage") or getattr(model_response, "usage", None) is None:
bedrock_input_tokens = response.headers.get("x-amzn-bedrock-input-token-count", None)
bedrock_output_tokens = response.headers.get("x-amzn-bedrock-output-token-count", None)
prompt_tokens = int(bedrock_input_tokens or litellm.token_counter(messages=messages))
completion_tokens = int(
bedrock_output_tokens
or litellm.token_counter(
text=model_response.choices[0].message.content, # type: ignore
count_response_tokens=True,
)
)
model_response.created = int(time.time())
model_response.model = model
usage = Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
)
setattr(model_response, "usage", usage)
else:
# Ensure created and model are set even if usage was already set
model_response.created = int(time.time())
model_response.model = model
return model_response
def completion(
self,
model: str,
messages: list,
api_base: Optional[str],
custom_prompt_dict: dict,
model_response: ModelResponse,
print_verbose: Callable,
encoding,
logging_obj: Logging,
optional_params: dict,
acompletion: bool,
timeout: Optional[Union[float, httpx.Timeout]],
litellm_params=None,
logger_fn=None,
extra_headers: Optional[dict] = None,
client: Optional[Union[AsyncHTTPHandler, HTTPHandler]] = None,
) -> Union[ModelResponse, CustomStreamWrapper]:
try:
from botocore.credentials import Credentials
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
## SETUP ##
stream = optional_params.pop("stream", None)
stream_chunk_size = optional_params.pop("stream_chunk_size", None)
provider = self.get_bedrock_invoke_provider(model)
modelId = self.get_bedrock_model_id(
model=model,
provider=provider,
optional_params=optional_params,
)
## CREDENTIALS ##
# pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them
aws_secret_access_key = optional_params.pop("aws_secret_access_key", None)
aws_access_key_id = optional_params.pop("aws_access_key_id", None)
aws_session_token = optional_params.pop("aws_session_token", None)
aws_region_name = optional_params.pop("aws_region_name", None)
aws_role_name = optional_params.pop("aws_role_name", None)
aws_session_name = optional_params.pop("aws_session_name", None)
aws_profile_name = optional_params.pop("aws_profile_name", None)
aws_bedrock_runtime_endpoint = optional_params.pop(
"aws_bedrock_runtime_endpoint", None
) # https://bedrock-runtime.{region_name}.amazonaws.com
aws_web_identity_token = optional_params.pop("aws_web_identity_token", None)
aws_sts_endpoint = optional_params.pop("aws_sts_endpoint", None)
ssl_verify = optional_params.pop("ssl_verify", None)
### SET REGION NAME ###
if aws_region_name is None:
# check env #
litellm_aws_region_name = get_secret("AWS_REGION_NAME", None)
if litellm_aws_region_name is not None and isinstance(litellm_aws_region_name, str):
aws_region_name = litellm_aws_region_name
standard_aws_region_name = get_secret("AWS_REGION", None)
if standard_aws_region_name is not None and isinstance(standard_aws_region_name, str):
aws_region_name = standard_aws_region_name
if aws_region_name is None:
aws_region_name = "us-west-2"
credentials: Credentials = self.get_credentials(
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
aws_region_name=aws_region_name,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
ssl_verify=ssl_verify,
)
### SET RUNTIME ENDPOINT ###
endpoint_url, proxy_endpoint_url = self.get_runtime_endpoint(
api_base=api_base,
aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint,
aws_region_name=aws_region_name,
)
if (stream is not None and stream is True) and provider != "ai21":
endpoint_url = f"{endpoint_url}/model/{modelId}/invoke-with-response-stream"
proxy_endpoint_url = f"{proxy_endpoint_url}/model/{modelId}/invoke-with-response-stream"
else:
endpoint_url = f"{endpoint_url}/model/{modelId}/invoke"
proxy_endpoint_url = f"{proxy_endpoint_url}/model/{modelId}/invoke"
if acompletion and provider == "anthropic" and self.is_claude_messages_api_model(model):
if isinstance(client, HTTPHandler):
client = None
return self._async_anthropic_messages_completion(
model=model,
messages=messages,
endpoint_url=endpoint_url,
proxy_endpoint_url=proxy_endpoint_url,
credentials=credentials,
aws_region_name=aws_region_name,
model_response=model_response,
print_verbose=print_verbose,
encoding=encoding,
logging_obj=logging_obj,
optional_params=optional_params,
stream=stream,
litellm_params=litellm_params,
logger_fn=logger_fn,
extra_headers=extra_headers,
timeout=timeout,
client=client,
stream_chunk_size=stream_chunk_size,
) # type: ignore[return-value]
prompt, chat_history = self.convert_messages_to_prompt(model, messages, provider, custom_prompt_dict)
inference_params = copy.deepcopy(optional_params)
json_schemas: dict = {}
if provider == "cohere":
if model.startswith("cohere.command-r"):
## LOAD CONFIG
config = litellm.AmazonCohereChatConfig().get_config()
for k, v in config.items():
if (
k not in inference_params
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
inference_params[k] = v
_data = {"message": prompt, **inference_params}
if chat_history is not None:
_data["chat_history"] = chat_history
data = json.dumps(_data)
else:
## LOAD CONFIG
config = litellm.AmazonCohereConfig.get_config()
for k, v in config.items():
if (
k not in inference_params
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
inference_params[k] = v
if stream is True:
inference_params["stream"] = True # cohere requires stream = True in inference params
data = json.dumps({"prompt": prompt, **inference_params})
elif provider == "anthropic":
if self.is_claude_messages_api_model(model):
# Separate system prompt from rest of message
system_prompt_idx: list[int] = []
system_messages: list[str] = []
for idx, message in enumerate(messages):
if message["role"] == "system":
system_messages.append(message["content"])
system_prompt_idx.append(idx)
if len(system_prompt_idx) > 0:
inference_params["system"] = "\n".join(system_messages)
messages = [i for j, i in enumerate(messages) if j not in system_prompt_idx]
# Format rest of message according to anthropic guidelines
messages = prompt_factory(model=model, messages=messages, custom_llm_provider="anthropic_xml") # type: ignore
## LOAD CONFIG
config = litellm.AmazonAnthropicClaudeConfig.get_config()
for k, v in config.items():
if (
k not in inference_params
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
inference_params[k] = v
## Handle Tool Calling
if "tools" in inference_params:
_is_function_call = True
for tool in inference_params["tools"]:
json_schemas[tool["function"]["name"]] = tool["function"].get("parameters", None)
tool_calling_system_prompt = construct_tool_use_system_prompt(tools=inference_params["tools"])
inference_params["system"] = (
inference_params.get("system", "\n") + tool_calling_system_prompt
) # add the anthropic tool calling prompt to the system prompt
inference_params.pop("tools")
data = json.dumps({"messages": messages, **inference_params})
else:
## LOAD CONFIG
config = litellm.AmazonAnthropicConfig.get_config()
for k, v in config.items():
if (
k not in inference_params
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
inference_params[k] = v
data = json.dumps({"prompt": prompt, **inference_params})
elif provider == "ai21":
## LOAD CONFIG
config = litellm.AmazonAI21Config.get_config()
for k, v in config.items():
if (
k not in inference_params
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
inference_params[k] = v
data = json.dumps({"prompt": prompt, **inference_params})
elif provider == "mistral":
## LOAD CONFIG
config = litellm.AmazonMistralConfig.get_config()
for k, v in config.items():
if (
k not in inference_params
): # completion(top_k=3) > amazon_config(top_k=3) <- allows for dynamic variables to be passed in
inference_params[k] = v
data = json.dumps({"prompt": prompt, **inference_params})
elif provider == "amazon": # amazon titan
## LOAD CONFIG
config = litellm.AmazonTitanConfig.get_config()
for k, v in config.items():
if (
k not in inference_params
): # completion(top_k=3) > amazon_config(top_k=3) <- allows for dynamic variables to be passed in
inference_params[k] = v
data = json.dumps(
{
"inputText": prompt,
"textGenerationConfig": inference_params,
}
)
elif provider == "meta" or provider == "llama":
## LOAD CONFIG
config = litellm.AmazonLlamaConfig.get_config()
for k, v in config.items():
if (
k not in inference_params
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
inference_params[k] = v
data = json.dumps({"prompt": prompt, **inference_params})
elif provider == "openai":
## OpenAI imported models use OpenAI Chat Completions format (messages-based)
# Use AmazonBedrockOpenAIConfig for proper OpenAI transformation
openai_config = AmazonBedrockOpenAIConfig()
supported_params = openai_config.get_supported_openai_params(model=model)
# Filter to only supported OpenAI params
filtered_params = {k: v for k, v in inference_params.items() if k in supported_params}
# OpenAI uses messages format, not prompt
data = json.dumps({"messages": messages, **filtered_params})
else:
## LOGGING
logging_obj.pre_call(
input=messages,
api_key="",
additional_args={
"complete_input_dict": inference_params,
},
)
raise BedrockError(
status_code=404,
message="Bedrock Invoke HTTPX: Unknown provider={}, model={}. Try calling via converse route - `bedrock/converse/<model>`.".format(
provider, model
),
)
## COMPLETION CALL
headers = {"Content-Type": "application/json"}
if extra_headers is not None:
headers = {"Content-Type": "application/json", **extra_headers}
prepped = self.get_request_headers(
credentials=credentials,
aws_region_name=aws_region_name,
extra_headers=extra_headers,
endpoint_url=endpoint_url,
data=data,
headers=headers,
)
## LOGGING
logging_obj.pre_call(
input=messages,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": proxy_endpoint_url,
"headers": prepped.headers,
},
)
### ROUTING (ASYNC, STREAMING, SYNC)
if acompletion:
if isinstance(client, HTTPHandler):
client = None
if stream is True and provider != "ai21":
return self.async_streaming(
model=model,
messages=messages,
data=data,
api_base=proxy_endpoint_url,
model_response=model_response,
print_verbose=print_verbose,
encoding=encoding,
logging_obj=logging_obj,
optional_params=optional_params,
stream=True,
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=prepped.headers,
timeout=timeout,
client=client,
stream_chunk_size=stream_chunk_size,
) # type: ignore
### ASYNC COMPLETION
return self.async_completion(
model=model,
messages=messages,
data=data,
api_base=proxy_endpoint_url,
model_response=model_response,
print_verbose=print_verbose,
encoding=encoding,
logging_obj=logging_obj,
optional_params=optional_params,
stream=stream, # type: ignore
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=prepped.headers,
timeout=timeout,
client=client,
) # type: ignore
if client is None or isinstance(client, AsyncHTTPHandler):
_params = {}
if timeout is not None:
if isinstance(timeout, float) or isinstance(timeout, int):
timeout = httpx.Timeout(timeout)
_params["timeout"] = timeout
self.client = _get_httpx_client(_params) # type: ignore
else:
self.client = client
if (stream is not None and stream is True) and provider != "ai21":
response = self.client.post(
url=proxy_endpoint_url,
headers=prepped.headers, # type: ignore
data=data,
stream=stream,
logging_obj=logging_obj,
)
if response.status_code != 200:
raise BedrockError(status_code=response.status_code, message=str(response.read()))
decoder = AWSEventStreamDecoder(model=model)
completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size))
streaming_response = CustomStreamWrapper(
completion_stream=completion_stream,
model=model,
custom_llm_provider="bedrock",
logging_obj=logging_obj,
)
## LOGGING
logging_obj.post_call(
input=messages,
api_key="",
original_response=streaming_response,
additional_args={"complete_input_dict": data},
)
return streaming_response
try:
response = self.client.post(
url=proxy_endpoint_url,
headers=dict(prepped.headers),
data=data,
logging_obj=logging_obj,
)
response.raise_for_status()
except httpx.HTTPStatusError as err:
error_code = err.response.status_code
raise BedrockError(status_code=error_code, message=err.response.text)
except httpx.TimeoutException:
raise BedrockError(status_code=408, message="Timeout error occurred.")
return self.process_response(
model=model,
response=response,
model_response=model_response,
stream=stream,
logging_obj=logging_obj,
optional_params=optional_params,
api_key="",
data=data,
messages=messages,
print_verbose=print_verbose,
encoding=encoding,
)
async def _async_anthropic_messages_completion(
self,
model: str,
messages: list,
endpoint_url: str,
proxy_endpoint_url: str,
credentials,
aws_region_name: str,
model_response: ModelResponse,
print_verbose: Callable,
encoding,
logging_obj: Logging,
optional_params: dict,
stream,
litellm_params=None,
logger_fn=None,
extra_headers: Optional[dict] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[AsyncHTTPHandler] = None,
stream_chunk_size: Optional[int] = None,
) -> Union[ModelResponse, CustomStreamWrapper]:
transformed_request = await litellm.AmazonAnthropicClaudeConfig().async_transform_request(
model=model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params or {},
headers=extra_headers or {},
)
data = json.dumps(transformed_request)
headers = {"Content-Type": "application/json"}
if extra_headers is not None:
headers = {"Content-Type": "application/json", **extra_headers}
prepped = self.get_request_headers(
credentials=credentials,
aws_region_name=aws_region_name,
extra_headers=extra_headers,
endpoint_url=endpoint_url,
data=data,
headers=headers,
)
logging_obj.pre_call(
input=messages,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": proxy_endpoint_url,
"headers": prepped.headers,
},
)
if stream is True:
return await self.async_streaming(
model=model,
messages=messages,
data=data,
api_base=proxy_endpoint_url,
model_response=model_response,
print_verbose=print_verbose,
encoding=encoding,
logging_obj=logging_obj,
optional_params=optional_params,
stream=True,
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=prepped.headers,
timeout=timeout,
client=client,
stream_chunk_size=stream_chunk_size,
)
return await self.async_completion(
model=model,
messages=messages,
data=data,
api_base=proxy_endpoint_url,
model_response=model_response,
print_verbose=print_verbose,
encoding=encoding,
logging_obj=logging_obj,
optional_params=optional_params,
stream=stream, # type: ignore
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=prepped.headers,
timeout=timeout,
client=client,
)
async def async_completion(
self,
model: str,
messages: list,
api_base: str,
model_response: ModelResponse,
print_verbose: Callable,
data: str,
timeout: Optional[Union[float, httpx.Timeout]],
encoding,
logging_obj: Logging,
stream,
optional_params: dict,
litellm_params=None,
logger_fn=None,
headers={},
client: Optional[AsyncHTTPHandler] = None,
) -> Union[ModelResponse, CustomStreamWrapper]:
if client is None:
_params = {}
if timeout is not None:
if isinstance(timeout, float) or isinstance(timeout, int):
timeout = httpx.Timeout(timeout)
_params["timeout"] = timeout
client = get_async_httpx_client(params=_params, llm_provider=litellm.LlmProviders.BEDROCK) # type: ignore
else:
client = client # type: ignore
try:
response = await client.post(
api_base,
headers=headers,
data=data,
timeout=timeout,
logging_obj=logging_obj,
)
response.raise_for_status()
except httpx.HTTPStatusError as err:
error_code = err.response.status_code
raise BedrockError(status_code=error_code, message=err.response.text)
except httpx.TimeoutException:
raise BedrockError(status_code=408, message="Timeout error occurred.")
return self.process_response(
model=model,
response=response,
model_response=model_response,
stream=stream if isinstance(stream, bool) else False,
logging_obj=logging_obj,
api_key="",
data=data,
messages=messages,
print_verbose=print_verbose,
optional_params=optional_params,
encoding=encoding,
)
@track_llm_api_timing() # for streaming, we need to instrument the function calling the wrapper
async def async_streaming(
self,
model: str,
messages: list,
api_base: str,
model_response: ModelResponse,
print_verbose: Callable,
data: str,
timeout: Optional[Union[float, httpx.Timeout]],
encoding,
logging_obj: Logging,
stream,
optional_params: dict,
litellm_params=None,
logger_fn=None,
headers={},
client: Optional[AsyncHTTPHandler] = None,
stream_chunk_size: Optional[int] = None,
) -> CustomStreamWrapper:
# The call is not made here; instead, we prepare the necessary objects for the stream.
streaming_response = CustomStreamWrapper(
completion_stream=None,
make_call=partial(
make_call,
client=client,
api_base=api_base,
headers=headers,
data=data, # type: ignore
model=model,
messages=messages,
logging_obj=logging_obj,
fake_stream=True if "ai21" in api_base else False,
stream_chunk_size=stream_chunk_size,
),
model=model,
custom_llm_provider="bedrock",
logging_obj=logging_obj,
)
return streaming_response
@staticmethod
def _get_provider_from_model_path(
model_path: str,
) -> Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL]:
"""
Helper function to get the provider from a model path with format: provider/model-name
Args:
model_path (str): The model path (e.g., 'llama/arn:aws:bedrock:us-east-1:086734376398:imported-model/r4c4kewx2s0n' or 'anthropic/model-name')
Returns:
Optional[str]: The provider name, or None if no valid provider found
"""
parts = model_path.split("/")
if len(parts) >= 1:
provider = parts[0]
if provider in get_args(litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL):
return cast(litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL, provider)
return None
class AWSEventStreamDecoder:
def __init__(self, model: str, json_mode: Optional[bool] = False) -> None:
from botocore.parsers import EventStreamJSONParser

View file

@ -1109,8 +1109,10 @@ def get_bedrock_chat_config(model: str):
Returns:
The appropriate Bedrock config class instance
"""
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
bedrock_route = BedrockModelInfo.get_bedrock_route(model)
bedrock_invoke_provider = litellm.BedrockLLM.get_bedrock_invoke_provider(model=model)
bedrock_invoke_provider = BaseAWSLLM.get_bedrock_invoke_provider(model=model)
base_model = BedrockModelInfo.get_base_model(model)
# Handle explicit routes first

View file

@ -1626,7 +1626,7 @@ class OpenAIFilesAPI(BaseLLM):
openai_client: AsyncOpenAI,
) -> OpenAIFileObject:
response = await openai_client.files.create(**create_file_data) # type: ignore[arg-type]
return OpenAIFileObject(**response.model_dump())
return OpenAIFileObject.model_validate(response.model_dump())
def create_file(
self,
@ -1662,7 +1662,7 @@ class OpenAIFilesAPI(BaseLLM):
create_file_data=create_file_data, openai_client=openai_client
)
response = cast(OpenAI, openai_client).files.create(**create_file_data) # type: ignore[arg-type]
return OpenAIFileObject(**response.model_dump())
return OpenAIFileObject.model_validate(response.model_dump())
async def afile_content(
self,
@ -1986,7 +1986,7 @@ class OpenAIBatchesAPI(BaseLLM):
openai_client: AsyncOpenAI,
) -> LiteLLMBatch:
response = await openai_client.batches.create(**create_batch_data) # type: ignore[arg-type]
return LiteLLMBatch(**response.model_dump())
return LiteLLMBatch.model_validate(response.model_dump())
def create_batch(
self,
@ -2023,7 +2023,7 @@ class OpenAIBatchesAPI(BaseLLM):
)
response = cast(OpenAI, openai_client).batches.create(**create_batch_data) # type: ignore[arg-type]
return LiteLLMBatch(**response.model_dump())
return LiteLLMBatch.model_validate(response.model_dump())
async def aretrieve_batch(
self,
@ -2032,7 +2032,7 @@ class OpenAIBatchesAPI(BaseLLM):
) -> LiteLLMBatch:
verbose_logger.debug("retrieving batch, args= %s", retrieve_batch_data)
response = await openai_client.batches.retrieve(**retrieve_batch_data) # type: ignore[arg-type]
return LiteLLMBatch(**response.model_dump())
return LiteLLMBatch.model_validate(response.model_dump())
def retrieve_batch(
self,
@ -2068,7 +2068,7 @@ class OpenAIBatchesAPI(BaseLLM):
retrieve_batch_data=retrieve_batch_data, openai_client=openai_client
)
response = cast(OpenAI, openai_client).batches.retrieve(**retrieve_batch_data) # type: ignore[arg-type]
return LiteLLMBatch(**response.model_dump())
return LiteLLMBatch.model_validate(response.model_dump())
async def acancel_batch(
self,
@ -2077,7 +2077,7 @@ class OpenAIBatchesAPI(BaseLLM):
) -> LiteLLMBatch:
verbose_logger.debug("async cancelling batch, args= %s", cancel_batch_data)
response = await openai_client.batches.cancel(**cancel_batch_data)
return LiteLLMBatch(**response.model_dump())
return LiteLLMBatch.model_validate(response.model_dump())
def cancel_batch(
self,
@ -2117,7 +2117,7 @@ class OpenAIBatchesAPI(BaseLLM):
if not isinstance(openai_client, OpenAI):
raise ValueError("OpenAI client is not an instance of OpenAI. Make sure you passed a sync OpenAI client.")
response = openai_client.batches.cancel(**cancel_batch_data)
return LiteLLMBatch(**response.model_dump())
return LiteLLMBatch.model_validate(response.model_dump())
async def alist_batches(
self,
@ -2477,9 +2477,9 @@ class OpenAIAssistantsAPI(BaseLLM):
response_obj: Optional[OpenAIMessage] = None
if getattr(thread_message, "status", None) is None:
thread_message.status = "completed"
response_obj = OpenAIMessage(**thread_message.dict())
response_obj = OpenAIMessage.model_validate(thread_message.dict())
else:
response_obj = OpenAIMessage(**thread_message.dict())
response_obj = OpenAIMessage.model_validate(thread_message.dict())
return response_obj
# fmt: off
@ -2556,9 +2556,9 @@ class OpenAIAssistantsAPI(BaseLLM):
response_obj: Optional[OpenAIMessage] = None
if getattr(thread_message, "status", None) is None:
thread_message.status = "completed"
response_obj = OpenAIMessage(**thread_message.dict())
response_obj = OpenAIMessage.model_validate(thread_message.dict())
else:
response_obj = OpenAIMessage(**thread_message.dict())
response_obj = OpenAIMessage.model_validate(thread_message.dict())
return response_obj
async def async_get_messages(

View file

@ -280,7 +280,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
raw_response_headers = dict(raw_response.headers)
processed_headers = process_response_headers(raw_response_headers)
try:
response = ResponsesAPIResponse(**raw_response_json)
response = ResponsesAPIResponse.model_validate(raw_response_json)
except Exception:
verbose_logger.debug(f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct")
response = ResponsesAPIResponse.model_construct(**raw_response_json)
@ -506,7 +506,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
raise OpenAIError(message=raw_response.text, status_code=raw_response.status_code)
raw_response_headers = dict(raw_response.headers)
processed_headers = process_response_headers(raw_response_headers)
response = ResponsesAPIResponse(**raw_response_json)
response = ResponsesAPIResponse.model_validate(raw_response_json)
response._hidden_params["additional_headers"] = processed_headers
response._hidden_params["headers"] = raw_response_headers
@ -588,7 +588,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
raw_response_headers = dict(raw_response.headers)
processed_headers = process_response_headers(raw_response_headers)
response = ResponsesAPIResponse(**raw_response_json)
response = ResponsesAPIResponse.model_validate(raw_response_json)
response._hidden_params["additional_headers"] = processed_headers
response._hidden_params["headers"] = raw_response_headers
@ -647,7 +647,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
processed_headers = process_response_headers(raw_response_headers)
try:
response = ResponsesAPIResponse(**raw_response_json)
response = ResponsesAPIResponse.model_validate(raw_response_json)
except Exception:
verbose_logger.debug(f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct")
response = ResponsesAPIResponse.model_construct(**raw_response_json)

View file

@ -5,7 +5,7 @@ Why separate file? Make it easy to see how transformation works
"""
import re
from typing import List, Optional, Tuple, Literal
from typing import List, Optional, Sequence, Tuple, Literal
from litellm.types.llms.openai import AllMessageValues
from litellm.types.llms.vertex_ai import CachedContentRequestBody
@ -152,6 +152,20 @@ def separate_cached_messages(
return cached_messages, non_cached_messages
def cached_messages_end_on_supported_turn(cached_messages: Sequence[AllMessageValues]) -> bool:
"""
The cachedContents API rejects contents ending on a model turn, which is how it
classifies both assistant messages and tool results, with HTTP 400
"Requests ending with a model turn are not supported". System messages are
extracted into system_instruction before contents are built, so the terminal
turn is the last non-system message.
"""
non_system_messages = tuple(message for message in cached_messages if message.get("role") != "system")
if not non_system_messages:
return bool(cached_messages)
return non_system_messages[-1].get("role") not in ("assistant", "tool", "function")
def transform_openai_messages_to_gemini_context_caching(
model: str,
messages: List[AllMessageValues],

View file

@ -22,6 +22,7 @@ from litellm.types.llms.vertex_ai import (
from ..common_utils import VertexAIError, get_vertex_base_url
from ..vertex_llm_base import VertexBase
from .transformation import (
cached_messages_end_on_supported_turn,
separate_cached_messages,
transform_openai_messages_to_gemini_context_caching,
)
@ -308,6 +309,14 @@ class ContextCachingEndpoints(VertexBase):
if len(cached_messages) == 0:
return messages, optional_params, None
if not cached_messages_end_on_supported_turn(cached_messages):
verbose_logger.debug(
"Vertex AI context caching: cached message block ends on a model turn once "
"system messages are extracted, which the cachedContents API rejects. "
"Skipping context caching."
)
return messages, optional_params, None
# Gemini requires a minimum of 1024 tokens for context caching.
# Skip caching if the cached content is too small to avoid API errors.
if not is_prompt_caching_valid_prompt(
@ -459,6 +468,14 @@ class ContextCachingEndpoints(VertexBase):
if len(cached_messages) == 0:
return messages, optional_params, None
if not cached_messages_end_on_supported_turn(cached_messages):
verbose_logger.debug(
"Vertex AI context caching: cached message block ends on a model turn once "
"system messages are extracted, which the cachedContents API rejects. "
"Skipping context caching."
)
return messages, optional_params, None
# Gemini requires a minimum of 1024 tokens for context caching.
# Skip caching if the cached content is too small to avoid API errors.
if not is_prompt_caching_valid_prompt(

View file

@ -207,7 +207,7 @@ from .llms.azure.chat.o_series_handler import AzureOpenAIO1ChatCompletion
from .llms.azure.completion.handler import AzureTextCompletion
from .llms.azure_ai.anthropic.handler import AzureAnthropicChatCompletion
from .llms.azure_ai.embed import AzureAIEmbedding
from .llms.bedrock.chat import BedrockConverseLLM, BedrockLLM
from .llms.bedrock.chat import BedrockConverseLLM
from .llms.bedrock.embed.embedding import BedrockEmbedding
from .llms.bedrock.image_edit.handler import BedrockImageEdit
from .llms.bedrock.image_generation.image_handler import BedrockImageGeneration

View file

@ -3454,22 +3454,16 @@
},
"azure_ai/gpt-5.4-mini": {
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_above_272k_tokens": 1.5e-07,
"cache_read_input_token_cost_priority": 1.5e-07,
"cache_read_input_token_cost_above_272k_tokens_priority": 3e-07,
"input_cost_per_token": 7.5e-07,
"input_cost_per_token_above_272k_tokens": 1.5e-06,
"input_cost_per_token_priority": 1.5e-06,
"input_cost_per_token_above_272k_tokens_priority": 3e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 400000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 4.5e-06,
"output_cost_per_token_above_272k_tokens": 6.75e-06,
"output_cost_per_token_priority": 9e-06,
"output_cost_per_token_above_272k_tokens_priority": 1.35e-05,
"source": "https://ai.azure.com/catalog/models/gpt-5.4-mini",
"supported_endpoints": [
"/v1/chat/completions",
@ -3500,22 +3494,16 @@
},
"azure_ai/gpt-5.4-mini-2026-03-17": {
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_above_272k_tokens": 1.5e-07,
"cache_read_input_token_cost_priority": 1.5e-07,
"cache_read_input_token_cost_above_272k_tokens_priority": 3e-07,
"input_cost_per_token": 7.5e-07,
"input_cost_per_token_above_272k_tokens": 1.5e-06,
"input_cost_per_token_priority": 1.5e-06,
"input_cost_per_token_above_272k_tokens_priority": 3e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 400000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 4.5e-06,
"output_cost_per_token_above_272k_tokens": 6.75e-06,
"output_cost_per_token_priority": 9e-06,
"output_cost_per_token_above_272k_tokens_priority": 1.35e-05,
"source": "https://ai.azure.com/catalog/models/gpt-5.4-mini",
"supported_endpoints": [
"/v1/chat/completions",
@ -3546,22 +3534,16 @@
},
"azure_ai/gpt-5.4-nano": {
"cache_read_input_token_cost": 2e-08,
"cache_read_input_token_cost_above_272k_tokens": 4e-08,
"cache_read_input_token_cost_priority": 4e-08,
"cache_read_input_token_cost_above_272k_tokens_priority": 8e-08,
"input_cost_per_token": 2e-07,
"input_cost_per_token_above_272k_tokens": 4e-07,
"input_cost_per_token_priority": 4e-07,
"input_cost_per_token_above_272k_tokens_priority": 8e-07,
"litellm_provider": "azure_ai",
"max_input_tokens": 400000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.25e-06,
"output_cost_per_token_above_272k_tokens": 1.875e-06,
"output_cost_per_token_priority": 2.5e-06,
"output_cost_per_token_above_272k_tokens_priority": 3.75e-06,
"source": "https://ai.azure.com/catalog/models/gpt-5.4-nano",
"supported_endpoints": [
"/v1/chat/completions",
@ -3592,22 +3574,16 @@
},
"azure_ai/gpt-5.4-nano-2026-03-17": {
"cache_read_input_token_cost": 2e-08,
"cache_read_input_token_cost_above_272k_tokens": 4e-08,
"cache_read_input_token_cost_priority": 4e-08,
"cache_read_input_token_cost_above_272k_tokens_priority": 8e-08,
"input_cost_per_token": 2e-07,
"input_cost_per_token_above_272k_tokens": 4e-07,
"input_cost_per_token_priority": 4e-07,
"input_cost_per_token_above_272k_tokens_priority": 8e-07,
"litellm_provider": "azure_ai",
"max_input_tokens": 400000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.25e-06,
"output_cost_per_token_above_272k_tokens": 1.875e-06,
"output_cost_per_token_priority": 2.5e-06,
"output_cost_per_token_above_272k_tokens_priority": 3.75e-06,
"source": "https://ai.azure.com/catalog/models/gpt-5.4-nano",
"supported_endpoints": [
"/v1/chat/completions",
@ -7201,7 +7177,7 @@
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_token": 7.5e-07,
"litellm_provider": "azure",
"max_input_tokens": 1050000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
@ -7236,7 +7212,7 @@
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_token": 7.5e-07,
"litellm_provider": "azure",
"max_input_tokens": 1050000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
@ -7271,7 +7247,7 @@
"cache_read_input_token_cost": 2e-08,
"input_cost_per_token": 2e-07,
"litellm_provider": "azure",
"max_input_tokens": 1050000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
@ -7306,7 +7282,7 @@
"cache_read_input_token_cost": 2e-08,
"input_cost_per_token": 2e-07,
"litellm_provider": "azure",
"max_input_tokens": 1050000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",

View file

@ -5322,7 +5322,7 @@ class MCPServerManager:
]
}
)
db_mcp_servers = [LiteLLM_MCPServerTable(**r.model_dump()) for r in raw_rows]
db_mcp_servers = [LiteLLM_MCPServerTable.model_validate(r.model_dump()) for r in raw_rows]
verbose_logger.info(f"Found {len(db_mcp_servers)} MCP servers in database")
previous_registry = self.registry

View file

@ -2434,7 +2434,7 @@ class ExperimentalUIJWTToken:
if decrypted_token is None:
return None
try:
return UserAPIKeyAuth(**json.loads(decrypted_token))
return UserAPIKeyAuth.model_validate(json.loads(decrypted_token))
except Exception as e:
raise Exception(f"Invalid hash key. Hash key={hashed_token}. Decrypted token={decrypted_token}. Error: {e}")
@ -2553,7 +2553,7 @@ async def get_key_object(
code=status.HTTP_401_UNAUTHORIZED,
)
_response = UserAPIKeyAuth(**_valid_token.model_dump(exclude_none=True))
_response = UserAPIKeyAuth.model_validate(_valid_token.model_dump(exclude_none=True))
# Load object_permission if object_permission_id exists but object_permission is not loaded
if _response.object_permission_id and not _response.object_permission:

View file

@ -1,4 +1,5 @@
from typing import Any, Dict, FrozenSet
from collections.abc import Mapping
from typing import Dict, FrozenSet, List, Union
from fastapi import Request
@ -83,21 +84,17 @@ async def handle_oauth2_proxy_request(request: Request) -> UserAPIKeyAuth:
"(signature-validated) instead of header-trust."
)
auth_data: Dict[str, Any] = {}
for key, header in oauth2_config_mappings.items():
value = request.headers.get(header)
if not value:
continue
if key == "models":
auth_data[key] = [model.strip() for model in value.split(",")]
else:
auth_data[key] = value
auth_data: Mapping[str, Union[str, List[str]]] = {
key: [model.strip() for model in value.split(",")] if key == "models" else value
for key, header in oauth2_config_mappings.items()
if (value := request.headers.get(header))
}
verbose_proxy_logger.debug(
"Auth data before creating UserAPIKeyAuth object: keys=%s",
list(auth_data.keys()),
)
user_api_key_auth = UserAPIKeyAuth(**auth_data)
user_api_key_auth = UserAPIKeyAuth.model_validate(auth_data)
verbose_proxy_logger.debug(
"UserAPIKeyAuth object created with keys: %s",
list(user_api_key_auth.__fields_set__),

View file

@ -118,7 +118,7 @@ class IdentityStore:
if from_db is None:
raise KeyNotFoundError(hashed_token)
key = UserAPIKeyAuth(**from_db.model_dump(exclude_none=True))
key = UserAPIKeyAuth.model_validate(from_db.model_dump(exclude_none=True))
if key.object_permission_id and not key.object_permission:
try:

View file

@ -216,7 +216,7 @@ async def _user_has_admin_privileges(
teams = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": user_obj.teams}})
for team in teams:
team_obj = LiteLLM_TeamTable(**team.model_dump())
team_obj = LiteLLM_TeamTable.model_validate(team.model_dump())
if _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj):
return True
@ -288,7 +288,7 @@ async def _team_admin_can_invite_user(
for team in teams
if _is_user_team_admin(
user_api_key_dict=user_api_key_dict,
team_obj=LiteLLM_TeamTable(**team.model_dump()),
team_obj=LiteLLM_TeamTable.model_validate(team.model_dump()),
)
]
if not admin_team_ids:

View file

@ -459,7 +459,7 @@ if MCP_AVAILABLE:
payload_dict: dict[str, Any] = loaded
try:
return MCPServer(**payload_dict)
return MCPServer.model_validate(payload_dict)
except Exception as e:
verbose_proxy_logger.debug(f"Invalid temporary MCP server payload in Redis cache: {str(e)}")
return None
@ -704,7 +704,7 @@ if MCP_AVAILABLE:
except AttributeError:
payload_dict = payload.dict() # type: ignore[attr-defined]
payload_dict["credentials"] = inherited_credentials
return NewMCPServerRequest(**payload_dict)
return NewMCPServerRequest.model_validate(payload_dict)
def _build_temporary_mcp_server_record(
payload: NewMCPServerRequest,

View file

@ -308,7 +308,7 @@ async def add_new_member(
)
await _append_team_id_if_absent(prisma_client, new_member.user_id, team_id)
if _returned_user is not None:
returned_user = LiteLLM_UserTable(**_returned_user.model_dump())
returned_user = LiteLLM_UserTable.model_validate(_returned_user.model_dump())
elif new_member.user_email is not None:
new_user_defaults = get_new_internal_user_defaults(user_id=str(uuid.uuid4()), user_email=new_member.user_email)
## user email is not unique acc. to prisma schema -> future improvement
@ -323,11 +323,11 @@ async def add_new_member(
_returned_user = await prisma_client.insert_data(data=new_user_defaults, table_name="user") # type: ignore
if _returned_user is not None:
returned_user = LiteLLM_UserTable(**_returned_user.model_dump())
returned_user = LiteLLM_UserTable.model_validate(_returned_user.model_dump())
elif len(existing_user_row) == 1:
user_info = existing_user_row[0]
await _append_team_id_if_absent(prisma_client, user_info.user_id, team_id)
returned_user = LiteLLM_UserTable(**user_info.model_dump())
returned_user = LiteLLM_UserTable.model_validate(user_info.model_dump())
elif len(existing_user_row) > 1:
raise HTTPException(
status_code=400,
@ -354,7 +354,7 @@ async def add_new_member(
include={"litellm_budget_table": True},
)
returned_team_membership = LiteLLM_TeamMembership(**_returned_team_membership.model_dump())
returned_team_membership = LiteLLM_TeamMembership.model_validate(_returned_team_membership.model_dump())
if returned_user is None:
raise Exception("Unable to update user table with membership information!")

View file

@ -5398,7 +5398,7 @@ class ProxyConfig:
# decrypt values
for k, v in _litellm_params.items():
_litellm_params[k] = self._resolve_db_litellm_param(key=k, value=v)
_litellm_params = LiteLLM_Params(**_litellm_params)
_litellm_params = LiteLLM_Params.model_validate(_litellm_params)
else:
verbose_proxy_logger.error(
@ -5429,7 +5429,7 @@ class ProxyConfig:
# decrypt values
for k, v in _litellm_params.items():
_litellm_params[k] = self._resolve_db_litellm_param(key=k, value=v)
_litellm_params = LiteLLM_Params(**_litellm_params)
_litellm_params = LiteLLM_Params.model_validate(_litellm_params)
else:
verbose_proxy_logger.error(
f"Invalid model added to proxy db. Invalid litellm params. litellm_params={_litellm_params}"
@ -13063,7 +13063,7 @@ def _get_model_group_info(
_model_group_info = llm_router.get_model_group_info(model_group=model)
if _model_group_info is not None:
model_groups.append(ModelGroupInfoProxy(**_model_group_info.model_dump()))
model_groups.append(ModelGroupInfoProxy.model_validate(_model_group_info.model_dump()))
else:
model_group_info = ModelGroupInfoProxy(
model_group=model,
@ -14782,7 +14782,7 @@ async def update_config_general_settings(
)
try:
ConfigGeneralSettings(**{data.field_name: data.field_value})
ConfigGeneralSettings.model_validate({data.field_name: data.field_value})
except Exception:
raise HTTPException(
status_code=400,

View file

@ -3,38 +3,58 @@ Base repository class with common functionality.
"""
from abc import ABC, abstractmethod
from typing import Any, Dict, Generic, List, Optional, Type, TypeVar
from collections.abc import Iterable, Mapping, Sequence
from typing import Any, Dict, Generic, List, Optional, Protocol, Tuple, Type, TypeVar, Union, runtime_checkable
from pydantic import BaseModel
T = TypeVar("T", bound=BaseModel)
def _record_to_dict(record: Any) -> Dict[str, Any]:
if isinstance(record, dict):
return record
if hasattr(record, "model_dump") and callable(record.model_dump):
@runtime_checkable
class SupportsModelDump(Protocol):
def model_dump(self) -> Dict[str, object]: ...
@runtime_checkable
class SupportsDict(Protocol):
def dict(self) -> Dict[str, object]: ...
DbRecord = Union[
Mapping[str, object],
SupportsModelDump,
SupportsDict,
Sequence[Tuple[str, object]],
]
def record_to_dict(record: DbRecord) -> Mapping[str, object]:
"""Project a database record into a mapping of column name to value."""
if isinstance(record, SupportsModelDump):
return record.model_dump()
if hasattr(record, "dict") and callable(record.dict):
if isinstance(record, SupportsDict):
return record.dict()
return dict(record)
if isinstance(record, Mapping):
return record
return {key: value for key, value in record}
class BaseRepository(ABC, Generic[T]):
"""Abstract base class for all repositories."""
def __init__(self, prisma_client: Any):
def __init__(self, prisma_client: Any): # any-ok: PrismaClient is an untyped runtime wrapper
self._prisma_client = prisma_client
@property
def prisma_client(self) -> Any:
def prisma_client(self) -> Any: # any-ok: PrismaClient is an untyped runtime wrapper
if self._prisma_client is None:
raise RuntimeError("No DB Connected. See - https://docs.litellm.ai/docs/proxy/virtual_keys")
return self._prisma_client
@property
@abstractmethod
def table(self) -> Any:
def table(self) -> Any: # any-ok: Prisma table actions are reached through the untyped client wrapper
"""Return the Prisma table for this repository."""
...
@ -44,21 +64,15 @@ class BaseRepository(ABC, Generic[T]):
"""Return the domain model class for this repository."""
...
def _to_model(self, record: Any) -> Optional[T]:
def _to_model(self, record: Optional[DbRecord]) -> Optional[T]:
"""Convert a database record to a domain model."""
if record is None:
return None
return self.model_class(**_record_to_dict(record))
return self.model_class.model_validate(record_to_dict(record))
def _to_model_list(self, records: List[Any]) -> List[T]:
def _to_model_list(self, records: Iterable[Optional[DbRecord]]) -> List[T]:
"""Convert a list of database records to domain models."""
result: List[T] = []
for r in records:
if r is not None:
model = self._to_model(r)
if model is not None:
result.append(model)
return result
return [model for record in records if record is not None and (model := self._to_model(record)) is not None]
async def find_by_id(self, id_value: str, id_field: str = "id") -> Optional[T]:
"""Find a record by its primary key."""

View file

@ -26,10 +26,8 @@ class OrganizationRepository(BaseRepository[LiteLLM_OrganizationTable]):
async def find_by_alias(self, organization_alias: str) -> Optional[LiteLLM_OrganizationTable]:
"""Find an organization by alias."""
records = await self.table.find_many(where={"organization_alias": organization_alias})
if records:
return self._to_model(records[0])
return None
organizations = await self.find_many(where={"organization_alias": organization_alias})
return organizations[0] if organizations else None
async def create_organization(
self,

View file

@ -24,15 +24,12 @@ class ProjectRepository(BaseRepository[LiteLLM_ProjectTable]):
async def find_by_alias(self, project_alias: str) -> Optional[LiteLLM_ProjectTable]:
"""Find a project by alias."""
records = await self.table.find_many(where={"project_alias": project_alias})
if records:
return self._to_model(records[0])
return None
projects = await self.find_many(where={"project_alias": project_alias})
return projects[0] if projects else None
async def find_by_team_id(self, team_id: str) -> List[LiteLLM_ProjectTable]:
"""Find all projects belonging to a team."""
records = await self.table.find_many(where={"team_id": team_id})
return self._to_model_list(records)
return await self.find_many(where={"team_id": team_id})
async def create_project(
self,

View file

@ -3,55 +3,59 @@ Team repository for database operations on LiteLLM_TeamTable.
"""
import json
from collections.abc import Mapping
from datetime import datetime
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Type
from pydantic import TypeAdapter
from litellm.models.team import LiteLLM_TeamTable, Member
from litellm.repositories.base_repository import BaseRepository
from litellm.repositories.base_repository import (
BaseRepository,
DbRecord,
record_to_dict,
)
if TYPE_CHECKING:
from prisma import Prisma
_MEMBERS_WITH_ROLES_ADAPTER = TypeAdapter(list[Member])
_JSON_ENCODED_TEAM_FIELDS = (
"metadata",
"model_spend",
"model_max_budget",
"router_settings",
"budget_limits",
"members_with_roles",
)
class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
"""Repository for team database operations."""
@property
def table(self) -> Any:
def table(self) -> Any: # any-ok: PrismaClient.db is an untyped runtime wrapper
return self.prisma_client.db.litellm_teamtable
@property
def deleted_table(self) -> Any:
def deleted_table(self) -> Any: # any-ok: PrismaClient.db is an untyped runtime wrapper
return self.prisma_client.db.litellm_deletedteamtable
@property
def model_class(self) -> Type[LiteLLM_TeamTable]:
return LiteLLM_TeamTable
def _to_model(self, record: Any) -> Optional[LiteLLM_TeamTable]:
def _to_model(self, record: Optional[DbRecord]) -> Optional[LiteLLM_TeamTable]:
"""Convert a database record to a Team model."""
if record is None:
return None
data = record.dict() if hasattr(record, "dict") else dict(record)
data = {
field: json.loads(value) if field in _JSON_ENCODED_TEAM_FIELDS and isinstance(value, str) else value
for field, value in record_to_dict(record).items()
}
json_fields = [
"metadata",
"model_spend",
"model_max_budget",
"router_settings",
"budget_limits",
"members_with_roles",
]
for field in json_fields:
if isinstance(data.get(field), str):
data[field] = json.loads(data[field])
return LiteLLM_TeamTable(**data)
return LiteLLM_TeamTable.model_validate(data)
async def get_members_with_roles_locked(self, tx: "Prisma", team_id: str) -> List[Member]:
"""Return the team's members_with_roles, locking the row FOR UPDATE.
@ -103,8 +107,8 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
organization_id: Optional[str] = None,
admins: Optional[List[str]] = None,
members: Optional[List[str]] = None,
members_with_roles: Optional[Dict[str, Any]] = None,
metadata: Optional[Dict[str, Any]] = None,
members_with_roles: Optional[Mapping[str, object]] = None,
metadata: Optional[Mapping[str, object]] = None,
max_budget: Optional[float] = None,
soft_budget: Optional[float] = None,
models: Optional[List[str]] = None,
@ -115,7 +119,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
object_permission_id: Optional[str] = None,
) -> LiteLLM_TeamTable:
"""Create a new team."""
data: Dict[str, Any] = {"team_id": team_id}
data: Dict[str, object] = {"team_id": team_id}
if team_alias is not None:
data["team_alias"] = team_alias
if organization_id is not None:
@ -154,8 +158,8 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
organization_id: Optional[str] = None,
admins: Optional[List[str]] = None,
members: Optional[List[str]] = None,
members_with_roles: Optional[Dict[str, Any]] = None,
metadata: Optional[Dict[str, Any]] = None,
members_with_roles: Optional[Mapping[str, object]] = None,
metadata: Optional[Mapping[str, object]] = None,
max_budget: Optional[float] = None,
soft_budget: Optional[float] = None,
models: Optional[List[str]] = None,
@ -167,7 +171,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
object_permission_id: Optional[str] = None,
) -> Optional[LiteLLM_TeamTable]:
"""Update a team."""
data: Dict[str, Any] = {}
data: Dict[str, object] = {}
if team_alias is not None:
data["team_alias"] = team_alias
if organization_id is not None:
@ -228,9 +232,9 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]):
return team
def _build_archive_data(self, team: LiteLLM_TeamTable) -> Dict[str, Any]:
def _build_archive_data(self, team: LiteLLM_TeamTable) -> Dict[str, object]:
"""Build archive data dict with only columns that exist in LiteLLM_DeletedTeamTable."""
data: Dict[str, Any] = {"team_id": team.team_id}
data: Dict[str, object] = {"team_id": team.team_id}
if team.team_alias is not None:
data["team_alias"] = team.team_alias
if team.organization_id is not None:

View file

@ -3,14 +3,18 @@ VerificationToken repository for database operations on LiteLLM_VerificationToke
"""
import json
from collections.abc import Iterator, Mapping
from collections.abc import Mapping
from datetime import datetime
from typing import TYPE_CHECKING, Any, Protocol
from typing import TYPE_CHECKING, Any
from litellm.models.verification_token import (
LiteLLM_VerificationToken,
)
from litellm.repositories.base_repository import BaseRepository
from litellm.repositories.base_repository import (
BaseRepository,
DbRecord,
record_to_dict,
)
if TYPE_CHECKING:
from prisma.models import (
@ -19,11 +23,17 @@ if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
class _DictConvertible(Protocol):
def dict(self) -> dict[str, object]: ...
def __iter__(self) -> Iterator[tuple[str, object]]: ...
_JSON_ENCODED_TOKEN_FIELDS = (
"aliases",
"config",
"permissions",
"metadata",
"model_spend",
"model_max_budget",
"router_settings",
"budget_limits",
"litellm_budget_table",
)
class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]):
@ -46,31 +56,21 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]):
def model_class(self) -> type[LiteLLM_VerificationToken]:
return LiteLLM_VerificationToken
def _to_model(self, record: _DictConvertible | None) -> LiteLLM_VerificationToken | None:
def _to_model(self, record: DbRecord | None) -> LiteLLM_VerificationToken | None:
"""Convert a database record to a VerificationToken model."""
if record is None:
return None
data = record.dict() if hasattr(record, "dict") else dict(record)
json_fields = [
"aliases",
"config",
"permissions",
"metadata",
"model_spend",
"model_max_budget",
"router_settings",
"budget_limits",
"litellm_budget_table",
]
for field in json_fields:
value = data.get(field)
if isinstance(value, str):
data[field] = json.loads(value)
if data.get("org_id") is None and data.get("organization_id") is not None:
data["org_id"] = data["organization_id"]
decoded = {
field: json.loads(value) if field in _JSON_ENCODED_TOKEN_FIELDS and isinstance(value, str) else value
for field, value in record_to_dict(record).items()
}
organization_id = decoded.get("organization_id")
data = (
decoded
if decoded.get("org_id") is not None or organization_id is None
else {**decoded, "org_id": organization_id}
)
return LiteLLM_VerificationToken.model_validate(data)

View file

@ -3454,22 +3454,16 @@
},
"azure_ai/gpt-5.4-mini": {
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_above_272k_tokens": 1.5e-07,
"cache_read_input_token_cost_priority": 1.5e-07,
"cache_read_input_token_cost_above_272k_tokens_priority": 3e-07,
"input_cost_per_token": 7.5e-07,
"input_cost_per_token_above_272k_tokens": 1.5e-06,
"input_cost_per_token_priority": 1.5e-06,
"input_cost_per_token_above_272k_tokens_priority": 3e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 400000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 4.5e-06,
"output_cost_per_token_above_272k_tokens": 6.75e-06,
"output_cost_per_token_priority": 9e-06,
"output_cost_per_token_above_272k_tokens_priority": 1.35e-05,
"source": "https://ai.azure.com/catalog/models/gpt-5.4-mini",
"supported_endpoints": [
"/v1/chat/completions",
@ -3500,22 +3494,16 @@
},
"azure_ai/gpt-5.4-mini-2026-03-17": {
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_above_272k_tokens": 1.5e-07,
"cache_read_input_token_cost_priority": 1.5e-07,
"cache_read_input_token_cost_above_272k_tokens_priority": 3e-07,
"input_cost_per_token": 7.5e-07,
"input_cost_per_token_above_272k_tokens": 1.5e-06,
"input_cost_per_token_priority": 1.5e-06,
"input_cost_per_token_above_272k_tokens_priority": 3e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 400000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 4.5e-06,
"output_cost_per_token_above_272k_tokens": 6.75e-06,
"output_cost_per_token_priority": 9e-06,
"output_cost_per_token_above_272k_tokens_priority": 1.35e-05,
"source": "https://ai.azure.com/catalog/models/gpt-5.4-mini",
"supported_endpoints": [
"/v1/chat/completions",
@ -3546,22 +3534,16 @@
},
"azure_ai/gpt-5.4-nano": {
"cache_read_input_token_cost": 2e-08,
"cache_read_input_token_cost_above_272k_tokens": 4e-08,
"cache_read_input_token_cost_priority": 4e-08,
"cache_read_input_token_cost_above_272k_tokens_priority": 8e-08,
"input_cost_per_token": 2e-07,
"input_cost_per_token_above_272k_tokens": 4e-07,
"input_cost_per_token_priority": 4e-07,
"input_cost_per_token_above_272k_tokens_priority": 8e-07,
"litellm_provider": "azure_ai",
"max_input_tokens": 400000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.25e-06,
"output_cost_per_token_above_272k_tokens": 1.875e-06,
"output_cost_per_token_priority": 2.5e-06,
"output_cost_per_token_above_272k_tokens_priority": 3.75e-06,
"source": "https://ai.azure.com/catalog/models/gpt-5.4-nano",
"supported_endpoints": [
"/v1/chat/completions",
@ -3592,22 +3574,16 @@
},
"azure_ai/gpt-5.4-nano-2026-03-17": {
"cache_read_input_token_cost": 2e-08,
"cache_read_input_token_cost_above_272k_tokens": 4e-08,
"cache_read_input_token_cost_priority": 4e-08,
"cache_read_input_token_cost_above_272k_tokens_priority": 8e-08,
"input_cost_per_token": 2e-07,
"input_cost_per_token_above_272k_tokens": 4e-07,
"input_cost_per_token_priority": 4e-07,
"input_cost_per_token_above_272k_tokens_priority": 8e-07,
"litellm_provider": "azure_ai",
"max_input_tokens": 400000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.25e-06,
"output_cost_per_token_above_272k_tokens": 1.875e-06,
"output_cost_per_token_priority": 2.5e-06,
"output_cost_per_token_above_272k_tokens_priority": 3.75e-06,
"source": "https://ai.azure.com/catalog/models/gpt-5.4-nano",
"supported_endpoints": [
"/v1/chat/completions",
@ -7201,7 +7177,7 @@
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_token": 7.5e-07,
"litellm_provider": "azure",
"max_input_tokens": 1050000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
@ -7236,7 +7212,7 @@
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_token": 7.5e-07,
"litellm_provider": "azure",
"max_input_tokens": 1050000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
@ -7271,7 +7247,7 @@
"cache_read_input_token_cost": 2e-08,
"input_cost_per_token": 2e-07,
"litellm_provider": "azure",
"max_input_tokens": 1050000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
@ -7306,7 +7282,7 @@
"cache_read_input_token_cost": 2e-08,
"input_cost_per_token": 2e-07,
"litellm_provider": "azure",
"max_input_tokens": 1050000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
@ -24364,7 +24340,7 @@
"input_cost_per_token_batches": 3.75e-07,
"input_cost_per_token_priority": 1.5e-06,
"litellm_provider": "openai",
"max_input_tokens": 1050000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
@ -24410,7 +24386,7 @@
"input_cost_per_token_batches": 3.75e-07,
"input_cost_per_token_priority": 1.5e-06,
"litellm_provider": "openai",
"max_input_tokens": 1050000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
@ -24454,7 +24430,7 @@
"input_cost_per_token_flex": 1e-07,
"input_cost_per_token_batches": 1e-07,
"litellm_provider": "openai",
"max_input_tokens": 1050000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
@ -24497,7 +24473,7 @@
"input_cost_per_token_flex": 1e-07,
"input_cost_per_token_batches": 1e-07,
"litellm_provider": "openai",
"max_input_tokens": 1050000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",

View file

@ -1,6 +1,6 @@
{
"ANN001": {
"limit": 3142
"limit": 3118
},
"ANN002": {
"limit": 69
@ -24,7 +24,7 @@
"limit": 130
},
"ANN401": {
"limit": 2013
"limit": 2010
},
"ASYNC230": {
"limit": 14
@ -33,7 +33,7 @@
"limit": 4
},
"B006": {
"limit": 190
"limit": 188
},
"B008": {
"limit": 505
@ -42,7 +42,7 @@
"limit": 84
},
"B010": {
"limit": 197
"limit": 194
},
"B018": {
"limit": 5
@ -60,7 +60,7 @@
"limit": 4
},
"BLE001": {
"limit": 2902
"limit": 2899
},
"C401": {
"limit": 11
@ -81,7 +81,7 @@
"limit": 4
},
"C901": {
"limit": 316
"limit": 314
},
"D419": {
"limit": 9
@ -135,7 +135,7 @@
"limit": 30
},
"PERF401": {
"limit": 143
"limit": 142
},
"PERF402": {
"limit": 9
@ -180,7 +180,7 @@
"limit": 34
},
"PLR1714": {
"limit": 265
"limit": 261
},
"PLR1730": {
"limit": 10
@ -189,7 +189,7 @@
"limit": 4
},
"PLW0127": {
"limit": 44
"limit": 43
},
"PLW0133": {
"limit": 4
@ -222,7 +222,7 @@
"limit": 38
},
"RET504": {
"limit": 717
"limit": 716
},
"RUF010": {
"limit": 874
@ -261,7 +261,7 @@
"limit": 24
},
"SIM101": {
"limit": 63
"limit": 61
},
"SIM102": {
"limit": 324
@ -273,7 +273,7 @@
"limit": 6
},
"SIM114": {
"limit": 113
"limit": 111
},
"SIM115": {
"limit": 5
@ -288,7 +288,7 @@
"limit": 4
},
"SIM210": {
"limit": 12
"limit": 11
},
"SIM211": {
"limit": 4
@ -309,22 +309,22 @@
"limit": 2652
},
"TRY002": {
"limit": 548
"limit": 547
},
"TRY004": {
"limit": 98
},
"TRY201": {
"limit": 422
"limit": 420
},
"TRY203": {
"limit": 122
"limit": 121
},
"TRY300": {
"limit": 881
"limit": 879
},
"UP006": {
"limit": 12143
"limit": 12138
},
"UP007": {
"limit": 2526
@ -348,7 +348,7 @@
"limit": 5
},
"UP032": {
"limit": 629
"limit": 626
},
"UP034": {
"limit": 4
@ -363,6 +363,6 @@
"limit": 105
},
"UP045": {
"limit": 17823
"limit": 17805
}
}

View file

@ -19,7 +19,8 @@ sys.path.insert(
import pytest
import litellm
from litellm.llms.azure.azure import get_azure_ad_token_from_oidc
from litellm.llms.bedrock.chat import BedrockConverseLLM, BedrockLLM
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.bedrock.chat import BedrockConverseLLM
from litellm.secret_managers.aws_secret_manager_v2 import AWSSecretsManagerV2
from litellm.secret_managers.main import (
get_secret,
@ -160,7 +161,7 @@ def test_oidc_circle_v1_with_amazon():
aws_role_name = "arn:aws:iam::335785316107:role/litellm-github-unit-tests-circleci-v1-assume-only"
aws_web_identity_token = "oidc/circleci/"
bllm = BedrockLLM()
bllm = BaseAWSLLM()
creds = bllm.get_credentials(
aws_region_name="ca-west-1",
aws_web_identity_token=aws_web_identity_token,

View file

@ -33,7 +33,7 @@ from litellm import (
completion_cost,
embedding,
)
from litellm.llms.bedrock.chat import BedrockLLM
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_tools_pt
from base_llm_unit_tests import BaseLLMChatTest, BaseAnthropicChatTest
@ -225,7 +225,7 @@ def bedrock_session_token_creds():
aws_region_name = os.environ["AWS_REGION_NAME"]
aws_session_token = os.environ.get("AWS_SESSION_TOKEN")
bllm = BedrockLLM()
bllm = BaseAWSLLM()
if aws_session_token is not None:
# For local testing
creds = bllm.get_credentials(
@ -3573,40 +3573,11 @@ def test_bedrock_openai_model_id_extraction():
print(f"✓ Model ID extracted and encoded: {model_id}")
def test_bedrock_openai_convert_messages_to_prompt():
"""
Test that convert_messages_to_prompt returns empty string for OpenAI models.
"""
from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM
bedrock_llm = BedrockLLM()
messages = [
{"role": "system", "content": "You are helpful"},
{"role": "user", "content": "Hello"},
]
prompt, chat_history = bedrock_llm.convert_messages_to_prompt(
model="test-model", messages=messages, provider="openai", custom_prompt_dict={}
def test_bedrock_openai_response_parsing():
from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import (
AmazonBedrockOpenAIConfig,
)
# OpenAI models use messages directly, no prompt conversion
assert prompt == ""
assert chat_history is None
print("✓ convert_messages_to_prompt returns empty for OpenAI")
def test_bedrock_openai_response_parsing():
"""
Test that OpenAI responses are correctly parsed.
"""
from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM
from litellm import ModelResponse
from unittest.mock import Mock
import json
bedrock_llm = BedrockLLM()
# Mock OpenAI-style response
openai_response = {
"choices": [
{
@ -3627,34 +3598,24 @@ def test_bedrock_openai_response_parsing():
mock_response.status_code = 200
mock_response.headers = {}
model_response = ModelResponse()
mock_logging = Mock()
result = bedrock_llm.process_response(
result = AmazonBedrockOpenAIConfig().transform_response(
model="openai/arn:aws:bedrock:us-east-1:123:imported-model/test",
response=mock_response,
model_response=model_response,
stream=False,
logging_obj=mock_logging,
optional_params={},
api_key="",
data={},
raw_response=mock_response,
model_response=ModelResponse(),
logging_obj=Mock(),
request_data={},
messages=[{"role": "user", "content": "What is the capital of France?"}],
print_verbose=lambda x: None,
optional_params={},
litellm_params={},
encoding=None,
)
# Verify response content
assert result.choices[0].message.content == "The capital of France is Paris."
assert result.choices[0].finish_reason == "stop"
# Verify usage
assert result.usage.prompt_tokens == 10
assert result.usage.completion_tokens == 8
assert result.usage.total_tokens == 18
print("✓ OpenAI response parsing works correctly")
def test_bedrock_openai_request_transformation():
"""
@ -3846,43 +3807,20 @@ def test_bedrock_openai_multiple_message_types():
def test_bedrock_openai_error_handling():
"""
Test that errors from OpenAI models are properly handled.
"""
from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM
from litellm import ModelResponse
from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import (
AmazonBedrockOpenAIConfig,
)
from litellm.llms.bedrock.common_utils import BedrockError
from unittest.mock import Mock
import json
bedrock_llm = BedrockLLM()
error = AmazonBedrockOpenAIConfig().get_error_class(
error_message="ValidationException: bad request",
status_code=422,
headers={},
)
# Mock error response
mock_response = Mock()
mock_response.json.side_effect = Exception("Invalid JSON")
mock_response.text = "Invalid response"
mock_response.status_code = 422
model_response = ModelResponse()
mock_logging = Mock()
with pytest.raises(BedrockError) as exc_info:
bedrock_llm.process_response(
model="openai/arn:aws:bedrock:us-east-1:123:imported-model/test",
response=mock_response,
model_response=model_response,
stream=False,
logging_obj=mock_logging,
optional_params={},
api_key="",
data={},
messages=[],
print_verbose=lambda x: None,
encoding=None,
)
assert exc_info.value.status_code == 422
print("✓ Error handling works correctly")
assert isinstance(error, BedrockError)
assert error.status_code == 422
assert "ValidationException: bad request" in str(error)
# ============================================================================

View file

@ -18,7 +18,9 @@ import json
import os
import sys
import httpx
import pytest
import respx
sys.path.insert(0, os.path.abspath("../../../.."))
@ -747,6 +749,98 @@ async def test_output_file_content_vertex_unified_file_id_extracts_gcs_uri(monke
assert captured["custom_llm_provider"] == "vertex_ai"
def _vertex_predictions_row(custom_id, prompt_tokens, completion_tokens):
return {
"request": {
"contents": [{"role": "user", "parts": [{"text": "hi"}]}],
"labels": {"litellm_custom_id": custom_id},
},
"status": "",
"response": {
"candidates": [
{
"content": {"role": "model", "parts": [{"text": "ok"}]},
"finishReason": "STOP",
}
],
"usageMetadata": {
"promptTokenCount": prompt_tokens,
"candidatesTokenCount": completion_tokens,
"totalTokenCount": prompt_tokens + completion_tokens,
},
"modelVersion": "gemini-3.6-flash",
},
"processed_time": "2026-07-30T00:00:00.000000+00:00",
}
@pytest.fixture
def respx_interceptable_httpx_client(monkeypatch):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
yield
litellm.in_memory_llm_clients_cache.flush_cache()
@pytest.mark.asyncio
@respx.mock
async def test_output_file_content_vertex_managed_uri_accepted_by_real_validation(respx_interceptable_httpx_client):
managed_output_uri = (
"gs://litellm-bucket/litellm-vertex-files/publishers/google/models/"
"gemini-3.6-flash/abc-123/prediction-model/predictions.jsonl"
)
rows = [
_vertex_predictions_row("request-1", 10, 5),
_vertex_predictions_row("request-2", 20, 10),
]
route = respx.get(url__regex=r"https://storage\.googleapis\.com/storage/v1/b/litellm-bucket/o/.*").mock(
return_value=httpx.Response(200, content=_vertex_jsonl(rows))
)
file_content = await bu._fetch_batch_output_file_content(
_batch(managed_output_uri),
custom_llm_provider="vertex_ai",
litellm_params={
"api_key": "test-token",
"vertex_project": "proj-1",
"vertex_location": "us-central1",
"gcs_bucket_name": "litellm-bucket",
},
)
result = bu._get_file_content_as_dictionary(file_content)
assert route.call_count == 1
request = route.calls.last.request
assert request.url.raw_path == (
b"/storage/v1/b/litellm-bucket/o/"
b"litellm-vertex-files%2Fpublishers%2Fgoogle%2Fmodels%2Fgemini-3.6-flash"
b"%2Fabc-123%2Fprediction-model%2Fpredictions.jsonl?alt=media"
)
assert [row["custom_id"] for row in result] == ["request-1", "request-2"]
assert all(row["response"]["status_code"] == 200 for row in result)
assert all(row["response"]["body"]["model"] == "gemini-3.6-flash" for row in result)
assert [row["response"]["body"]["usage"]["prompt_tokens"] for row in result] == [10, 20]
assert [row["response"]["body"]["usage"]["completion_tokens"] for row in result] == [5, 10]
@pytest.mark.asyncio
@respx.mock
async def test_output_file_content_vertex_foreign_bucket_rejected_by_real_validation():
with pytest.raises(Exception, match="does not match the configured storage bucket"):
await bu._fetch_batch_output_file_content(
_batch("gs://attacker-bucket/litellm-vertex-files/x/predictions.jsonl"),
custom_llm_provider="vertex_ai",
litellm_params={
"api_key": "test-token",
"vertex_project": "proj-1",
"vertex_location": "us-central1",
"gcs_bucket_name": "litellm-bucket",
},
)
assert respx.mock.calls.call_count == 0
@pytest.mark.asyncio
async def test_handle_completed_vertex_batch_computes_cost_usage_and_models(monkeypatch):
import litellm.files.main as files_main

View file

@ -8,14 +8,11 @@ sys.path.insert(
0, os.path.abspath("../../../../..")
) # Adds the parent directory to the system path
import litellm
from litellm.llms.bedrock.chat.invoke_handler import (
AWSEventStreamDecoder,
BedrockLLM,
make_call,
make_sync_call,
)
from litellm.llms.custom_httpx.http_handler import HTTPHandler
def test_transform_thinking_blocks_with_redacted_content():
@ -296,33 +293,3 @@ def test_make_sync_call_honors_explicit_stream_chunk_size():
response.iter_bytes.assert_called_once_with(chunk_size=2048)
def test_legacy_bedrock_llm_streaming_does_not_rechunk_by_default():
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.iter_bytes = MagicMock(return_value=iter([]))
client = HTTPHandler()
client.post = MagicMock(return_value=mock_response)
BedrockLLM().completion(
model="cohere.command-text-v14",
messages=[{"role": "user", "content": "hi"}],
api_base=None,
custom_prompt_dict={},
model_response=litellm.ModelResponse(),
print_verbose=lambda *args, **kwargs: None,
encoding=litellm.encoding,
logging_obj=MagicMock(),
optional_params={
"stream": True,
"aws_access_key_id": "fake",
"aws_secret_access_key": "fake",
"aws_region_name": "us-east-1",
},
acompletion=False,
timeout=None,
litellm_params={},
client=client,
)
mock_response.iter_bytes.assert_called_once_with(chunk_size=None)

View file

@ -1452,6 +1452,157 @@ class TestContextCachingEndpoints:
# Restart the patcher so teardown_method can stop it cleanly
self._token_check_patcher.start()
def _model_turn_final_messages(self, final_cached_role):
tool_call = {
"id": "call_abc123",
"type": "function",
"function": {"name": "get_weather", "arguments": '{"location": "Boston"}'},
}
cached_tail = {
"assistant": [],
"tool": [
{
"role": "tool",
"tool_call_id": "call_abc123",
"content": "72F and sunny",
"cache_control": {"type": "ephemeral"},
}
],
"system": [
{
"role": "system",
"content": "Tool results are authoritative.",
"cache_control": {"type": "ephemeral"},
}
],
}[final_cached_role]
return [
{
"role": "user",
"content": [
{
"type": "text",
"text": "Use the weather tool for every answer.",
"cache_control": {"type": "ephemeral"},
}
],
},
{
"role": "assistant",
"content": "",
"tool_calls": [tool_call],
"cache_control": {"type": "ephemeral"},
},
*cached_tail,
{"role": "user", "content": "What is the weather in Boston?"},
]
@pytest.mark.parametrize("final_cached_role", ["assistant", "tool", "system"])
def test_check_and_create_cache_skips_when_cached_block_ends_on_model_turn(
self, final_cached_role
):
"""The cachedContents API rejects contents ending on an assistant or tool turn
with HTTP 400 "Requests ending with a model turn are not supported", so the
request must proceed uncached instead of failing.
"""
all_messages = self._model_turn_final_messages(final_cached_role)
optional_params = self.sample_optional_params.copy()
result = self.context_caching.check_and_create_cache(
messages=all_messages,
optional_params=optional_params,
api_key="test_key",
api_base=None,
model="gemini-3.6-flash",
client=self.mock_client,
timeout=30.0,
logging_obj=self.mock_logging,
cached_content=None,
custom_llm_provider="vertex_ai",
vertex_project="test_project",
vertex_location="us-central1",
vertex_auth_header="test_token",
)
messages, returned_params, returned_cache = result
assert messages == all_messages
assert returned_cache is None
assert "tools" in returned_params
self.mock_client.get.assert_not_called()
self.mock_client.post.assert_not_called()
@pytest.mark.parametrize("final_cached_role", ["assistant", "tool", "system"])
@pytest.mark.asyncio
async def test_async_check_and_create_cache_skips_when_cached_block_ends_on_model_turn(
self, final_cached_role
):
"""Async variant: an unsupported terminal turn skips caching instead of failing."""
all_messages = self._model_turn_final_messages(final_cached_role)
optional_params = self.sample_optional_params.copy()
result = await self.context_caching.async_check_and_create_cache(
messages=all_messages,
optional_params=optional_params,
api_key="test_key",
api_base=None,
model="gemini-3.6-flash",
client=self.mock_async_client,
timeout=30.0,
logging_obj=self.mock_logging,
cached_content=None,
custom_llm_provider="vertex_ai",
vertex_project="test_project",
vertex_location="us-central1",
vertex_auth_header="test_token",
)
messages, returned_params, returned_cache = result
assert messages == all_messages
assert returned_cache is None
assert "tools" in returned_params
self.mock_async_client.get.assert_not_called()
self.mock_async_client.post.assert_not_called()
def test_cached_messages_end_on_supported_turn():
from litellm.llms.vertex_ai.context_caching.transformation import (
cached_messages_end_on_supported_turn,
)
assert (
cached_messages_end_on_supported_turn(
[{"role": "assistant", "content": "hi"}, {"role": "user", "content": "hello"}]
)
is True
)
assert cached_messages_end_on_supported_turn([{"role": "system", "content": "be brief"}]) is True
assert cached_messages_end_on_supported_turn([{"role": "assistant", "content": "hi"}]) is False
assert (
cached_messages_end_on_supported_turn(
[
{"role": "user", "content": "hello"},
{"role": "assistant", "content": "hi"},
{"role": "system", "content": "be brief"},
]
)
is False
)
assert (
cached_messages_end_on_supported_turn(
[{"role": "system", "content": "be brief"}, {"role": "user", "content": "hello"}]
)
is True
)
assert (
cached_messages_end_on_supported_turn([{"role": "tool", "tool_call_id": "x", "content": "y"}])
is False
)
assert (
cached_messages_end_on_supported_turn([{"role": "function", "name": "f", "content": "y"}])
is False
)
assert cached_messages_end_on_supported_turn([]) is False
class TestCheckCachePagination:
"""Test pagination logic in check_cache and async_check_cache methods."""

View file

@ -203,23 +203,32 @@ class TestBaseRepository:
assert len(budgets) == 1
def test_record_to_dict_branches(self):
from litellm.repositories.base_repository import _record_to_dict
from litellm.repositories.base_repository import record_to_dict
assert _record_to_dict({"a": 1}) == {"a": 1}
assert record_to_dict({"a": 1}) == {"a": 1}
class WithModelDump:
def model_dump(self):
return {"src": "model_dump"}
assert _record_to_dict(WithModelDump()) == {"src": "model_dump"}
assert record_to_dict(WithModelDump()) == {"src": "model_dump"}
class WithDict:
def dict(self):
return {"src": "dict"}
assert _record_to_dict(WithDict()) == {"src": "dict"}
assert record_to_dict(WithDict()) == {"src": "dict"}
assert _record_to_dict([("k", "v")]) == {"k": "v"}
assert record_to_dict([("k", "v")]) == {"k": "v"}
class WithBoth:
def model_dump(self):
return {"src": "model_dump"}
def dict(self):
return {"src": "dict"}
assert record_to_dict(WithBoth()) == {"src": "model_dump"}
class TestBudgetRepository:

View file

@ -0,0 +1,81 @@
import json
from functools import lru_cache
from pathlib import Path
import pytest
REPO_ROOT = Path(__file__).parents[2]
MAIN_PATH = REPO_ROOT / "model_prices_and_context_window.json"
BACKUP_PATH = REPO_ROOT / "litellm" / "model_prices_and_context_window_backup.json"
DOCUMENTED_MAX_INPUT_TOKENS = 272000
DOCUMENTED_MAX_OUTPUT_TOKENS = 128000
SMALL_MODEL_NAMES = (
"gpt-5.4-mini",
"gpt-5.4-mini-2026-03-17",
"gpt-5.4-nano",
"gpt-5.4-nano-2026-03-17",
)
SMALL_MODELS = tuple(f"{prefix}{name}" for prefix in ("", "azure/", "azure_ai/") for name in SMALL_MODEL_NAMES)
STANDARD_PRICING = {
"gpt-5.4-mini": (7.5e-07, 4.5e-06, 7.5e-08),
"gpt-5.4-nano": (2e-07, 1.25e-06, 2e-08),
}
LONG_CONTEXT_MODELS = ("gpt-5.4", "gpt-5.4-pro")
@lru_cache(maxsize=2)
def _load(path: Path) -> dict[str, dict[str, object]]:
with open(path) as f:
return json.load(f)
def _pricing_key(model: str) -> str:
return "gpt-5.4-nano" if "nano" in model else "gpt-5.4-mini"
@pytest.mark.parametrize("model", SMALL_MODELS)
def test_gpt_5_4_small_models_use_documented_token_limits(model: str) -> None:
"""gpt-5.4-mini/nano are 400K-window models: 272K in, 128K out, not gpt-5.4's 1.05M window."""
info = _load(MAIN_PATH).get(model)
assert info is not None, f"{model} not found in model_prices_and_context_window.json"
assert info["max_input_tokens"] == DOCUMENTED_MAX_INPUT_TOKENS
assert info["max_output_tokens"] == DOCUMENTED_MAX_OUTPUT_TOKENS
assert info["max_tokens"] == DOCUMENTED_MAX_OUTPUT_TOKENS
@pytest.mark.parametrize("model", SMALL_MODELS)
def test_gpt_5_4_small_models_have_no_long_context_surcharge(model: str) -> None:
"""OpenAI prices prompts above 272K at 2x input / 1.5x output for the 1.05M-window models only."""
info = _load(MAIN_PATH)[model]
assert [key for key in info if "above_272k" in key] == []
@pytest.mark.parametrize("model", SMALL_MODELS)
def test_gpt_5_4_small_models_standard_pricing(model: str) -> None:
info = _load(MAIN_PATH)[model]
input_cost, output_cost, cache_read_cost = STANDARD_PRICING[_pricing_key(model)]
assert info["input_cost_per_token"] == input_cost
assert info["output_cost_per_token"] == output_cost
assert info["cache_read_input_token_cost"] == cache_read_cost
@pytest.mark.parametrize("model", LONG_CONTEXT_MODELS)
def test_gpt_5_4_long_context_models_keep_surcharge(model: str) -> None:
"""The mini/nano correction must leave gpt-5.4 and gpt-5.4-pro tiered pricing intact."""
info = _load(MAIN_PATH)[model]
assert info["input_cost_per_token_above_272k_tokens"] == pytest.approx(info["input_cost_per_token"] * 2)
assert info["output_cost_per_token_above_272k_tokens"] == pytest.approx(info["output_cost_per_token"] * 1.5)
@pytest.mark.parametrize("model", SMALL_MODELS)
def test_gpt_5_4_small_models_backup_matches_main(model: str) -> None:
assert _load(BACKUP_PATH).get(model) == _load(MAIN_PATH).get(model), (
f"{model} differs between main and backup model cost maps"
)

View file

@ -17,7 +17,6 @@ sys.path.insert(0, str(Path(__file__).parent))
import litellm.proxy.guardrails.guardrail_hooks.aim.aim as _aim_module
import litellm.proxy.guardrails.guardrail_hooks.cato_networks.cato_networks as _cato_networks_module
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.bedrock.chat.invoke_handler import BedrockLLM
from litellm.proxy.guardrails.guardrail_hooks.aim.aim import AimGuardrail
from litellm.proxy.guardrails.guardrail_hooks.cato_networks.cato_networks import CatoNetworksGuardrail
@ -87,23 +86,6 @@ class TestBaseAWSLLMSSLVerify:
assert True # If we got here without error, parameter was accepted
class TestBedrockLLMSSLVerify:
"""Test SSL verification parameter handling in BedrockLLM."""
def test_bedrock_llm_accepts_ssl_verify_in_optional_params(self):
"""Test that BedrockLLM can receive ssl_verify in optional_params."""
# This is a simple test to verify the parameter is accepted
# The actual propagation is tested in integration tests
bedrock_llm = BedrockLLM()
# Verify the class exists and can be instantiated
assert bedrock_llm is not None
# Verify _get_ssl_verify method exists and works
result = bedrock_llm._get_ssl_verify(ssl_verify="/path/to/cert.pem")
assert result == "/path/to/cert.pem"
class TestAimGuardrailSSLVerify:
"""Test SSL verification parameter handling in AimGuardrail."""

View file

@ -1,9 +1,9 @@
{
"LIT001": {
"limit": 23279
"limit": 23253
},
"LIT002": {
"limit": 27473
"limit": 27433
},
"LIT003": {
"limit": 292
@ -15,7 +15,7 @@
"limit": 0
},
"LIT006": {
"limit": 1109
"limit": 1108
},
"LIT007": {
"limit": 0
@ -24,6 +24,6 @@
"limit": 1004
},
"LIT009": {
"limit": 2495
"limit": 2474
}
}