mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
commit
a97233067d
37 changed files with 1164 additions and 1336 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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": ""
|
||||
}
|
||||
|
|
@ -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.
|
||||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -5,7 +5,6 @@ from .invoke_handler import (
|
|||
AmazonAnthropicClaudeStreamDecoder,
|
||||
AmazonDeepSeekR1StreamDecoder,
|
||||
AWSEventStreamDecoder,
|
||||
BedrockLLM,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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__),
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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!")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
# ============================================================================
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
81
tests/test_litellm/test_gpt_5_4_model_metadata.py
Normal file
81
tests/test_litellm/test_gpt_5_4_model_metadata.py
Normal 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"
|
||||
)
|
||||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue