mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_rebuild_admin_ui_static_export
This commit is contained in:
commit
e4585b2fa4
927 changed files with 57459 additions and 8768 deletions
|
|
@ -158,6 +158,8 @@ jobs:
|
|||
CHOCOLATEY_CONFIRM_ALL: "true"
|
||||
- run:
|
||||
name: Install Dependencies
|
||||
environment:
|
||||
UV_HTTP_TIMEOUT: "300"
|
||||
command: |
|
||||
$installer = Join-Path $env:TEMP "uv-install.ps1"
|
||||
Invoke-WebRequest -Uri https://astral.sh/uv/0.10.9/install.ps1 -OutFile $installer
|
||||
|
|
@ -2475,10 +2477,15 @@ jobs:
|
|||
DISABLE_SCHEMA_UPDATE: "true"
|
||||
SERVER_ROOT_PATH: ""
|
||||
PROXY_LOGOUT_URL: ""
|
||||
# LITELLM_LICENSE is forwarded from the project env so premium-gated
|
||||
# UI flows can be exercised. license.spec.ts asserts the resulting
|
||||
# JWT carries premium_user=true; if it ever stops being passed, that
|
||||
# test fails loudly rather than silently regressing premium coverage.
|
||||
command: |
|
||||
uv run --no-sync python -m litellm.proxy.proxy_cli \
|
||||
--config ui/litellm-dashboard/e2e_tests/fixtures/config.yml \
|
||||
--port 4000
|
||||
LITELLM_LICENSE="$LITELLM_LICENSE" \
|
||||
uv run --no-sync python -m litellm.proxy.proxy_cli \
|
||||
--config ui/litellm-dashboard/e2e_tests/fixtures/config.yml \
|
||||
--port 4000
|
||||
background: true
|
||||
- run:
|
||||
name: Wait for proxy to be ready
|
||||
|
|
@ -2495,9 +2502,12 @@ jobs:
|
|||
exit 1
|
||||
- run:
|
||||
name: Run Playwright E2E tests
|
||||
# Forward LITELLM_LICENSE so license.spec.ts can detect that the
|
||||
# proxy was launched with a license and assert premium_user=true.
|
||||
command: |
|
||||
cd ui/litellm-dashboard
|
||||
npx playwright test --config e2e_tests/playwright.config.ts
|
||||
LITELLM_LICENSE="$LITELLM_LICENSE" \
|
||||
npx playwright test --config e2e_tests/playwright.config.ts
|
||||
no_output_timeout: 10m
|
||||
- store_artifacts:
|
||||
path: ui/litellm-dashboard/test-results
|
||||
|
|
@ -2531,7 +2541,6 @@ jobs:
|
|||
paths:
|
||||
- litellm-docker-database.tar.zst
|
||||
|
||||
|
||||
test_bad_database_url:
|
||||
machine:
|
||||
image: ubuntu-2204:2024.04.1
|
||||
|
|
|
|||
28
.github/workflows/codeql.yml
vendored
28
.github/workflows/codeql.yml
vendored
|
|
@ -53,3 +53,31 @@ jobs:
|
|||
uses: github/codeql-action/analyze@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3
|
||||
with:
|
||||
category: "/language:${{ matrix.language }}"
|
||||
output: sarif-results
|
||||
upload: failure-only
|
||||
|
||||
# py/weak-sensitive-data-hashing (CWE-328) fires on the OCI signing call at
|
||||
# litellm/llms/oci/common_utils.py, which hashes the HTTP request body to
|
||||
# produce the x-content-sha256 header required by the OCI HTTP signing spec —
|
||||
# a content-integrity hash, not a password or secret hash. SHA-256 is mandated
|
||||
# by Oracle for this header; see
|
||||
# https://docs.oracle.com/en-us/iaas/Content/API/Concepts/signingrequests.htm
|
||||
# The `usedforsecurity=False` flag on the hashlib.sha256 call already declares
|
||||
# non-security intent, but CodeQL's taint flow still re-fires when callers
|
||||
# further up the stack are modified. The suppression is scoped to this one
|
||||
# file/rule pair via SARIF post-filtering so every other callsite of
|
||||
# py/weak-sensitive-data-hashing in the repository continues to be analyzed.
|
||||
- name: Filter SARIF (OCI sha256)
|
||||
if: matrix.language == 'python'
|
||||
uses: advanced-security/filter-sarif@2da736ff05ef065cb2894ac6892e47b5eac2c3c0 # v1.1
|
||||
with:
|
||||
patterns: |
|
||||
-litellm/llms/oci/common_utils.py:py/weak-sensitive-data-hashing
|
||||
input: sarif-results/python.sarif
|
||||
output: sarif-results/python.sarif
|
||||
|
||||
- name: Upload SARIF
|
||||
uses: github/codeql-action/upload-sarif@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3
|
||||
with:
|
||||
sarif_file: sarif-results
|
||||
category: "/language:${{ matrix.language }}"
|
||||
|
|
|
|||
2
.github/workflows/test-unit-proxy-db.yml
vendored
2
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -215,8 +215,10 @@ jobs:
|
|||
tests/proxy_unit_tests/test_models_fallback_endpoint.py
|
||||
tests/proxy_unit_tests/test_google_endpoint_routing.py
|
||||
tests/proxy_unit_tests/test_google_gemini_proxy_request.py
|
||||
tests/proxy_unit_tests/test_gemini_agents_endpoints.py
|
||||
tests/proxy_unit_tests/test_get_favicon.py
|
||||
tests/proxy_unit_tests/test_get_image.py
|
||||
tests/proxy_unit_tests/test_reducto_ocr_route.py
|
||||
tests/proxy_unit_tests/test_ui_path_detection.py
|
||||
tests/proxy_unit_tests/test_prompt_test_endpoint.py
|
||||
tests/proxy_unit_tests/test_check_batch_cost.py
|
||||
|
|
|
|||
34
.github/workflows/test-unit-proxy-mgmt-behavior.yml
vendored
Normal file
34
.github/workflows/test-unit-proxy-mgmt-behavior.yml
vendored
Normal file
|
|
@ -0,0 +1,34 @@
|
|||
name: "Unit Tests: Proxy Management-Endpoint Behavior Pinning"
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_branch
|
||||
- "litellm_**"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
proxy-mgmt-behavior:
|
||||
uses: ./.github/workflows/_test-unit-services-base.yml
|
||||
with:
|
||||
test-path: tests/proxy_behavior
|
||||
# workers=0 (no xdist): the world seed is a single shared Postgres
|
||||
# state — two xdist workers both call seed_world() and race on the
|
||||
# ``behavior-pin-budget`` row, producing UniqueViolation + cascading
|
||||
# missing-membership FK failures. The whole suite is ~7s sequentially,
|
||||
# so the cost of disabling parallelism here is negligible.
|
||||
workers: 0
|
||||
reruns: 0
|
||||
enable-postgres: true
|
||||
artifact-name: proxy-mgmt-behavior
|
||||
timeout-minutes: 15
|
||||
|
|
@ -292,7 +292,7 @@ curl -X POST 'http://0.0.0.0:4000/v1/chat/completions' \
|
|||
| [CompactifAI (`compactifai`)](https://docs.litellm.ai/docs/providers/compactifai) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Custom (`custom`)](https://docs.litellm.ai/docs/providers/custom_llm_server) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Custom OpenAI (`custom_openai`)](https://docs.litellm.ai/docs/providers/openai_compatible) | ✅ | ✅ | ✅ | | | ✅ | ✅ | ✅ | ✅ | |
|
||||
| [Dashscope (`dashscope`)](https://docs.litellm.ai/docs/providers/dashscope) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Dashscope (`dashscope`)](https://docs.litellm.ai/docs/providers/dashscope) | ✅ | ✅ | ✅ | ✅ | | | | | | ✅ |
|
||||
| [Databricks (`databricks`)](https://docs.litellm.ai/docs/providers/databricks) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [DataRobot (`datarobot`)](https://docs.litellm.ai/docs/providers/datarobot) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Deepgram (`deepgram`)](https://docs.litellm.ai/docs/providers/deepgram) | ✅ | ✅ | ✅ | | | ✅ | | | | |
|
||||
|
|
|
|||
|
|
@ -24,7 +24,7 @@ version: 1.1.0
|
|||
# incremented each time you make changes to the application. Versions are not expected to
|
||||
# follow Semantic Versioning. They should reflect the version the application is using.
|
||||
# It is recommended to use it with quotes.
|
||||
appVersion: v1.80.12
|
||||
appVersion: v1.85.1
|
||||
|
||||
annotations:
|
||||
org.opencontainers.image.source: "https://github.com/BerriAI/litellm"
|
||||
|
|
|
|||
|
|
@ -53,7 +53,7 @@ spec:
|
|||
- name: {{ include "litellm.name" . }}
|
||||
securityContext:
|
||||
{{- toYaml .Values.securityContext | nindent 12 }}
|
||||
image: "{{ .Values.image.repository }}:{{ .Values.image.tag | default (printf "main-%s" .Chart.AppVersion) }}"
|
||||
image: "{{ .Values.image.repository }}:{{ .Values.image.tag | default .Chart.AppVersion }}"
|
||||
imagePullPolicy: {{ .Values.image.pullPolicy }}
|
||||
env:
|
||||
- name: HOST
|
||||
|
|
|
|||
|
|
@ -12,6 +12,10 @@ spec:
|
|||
name: {{ include "litellm.fullname" . }}
|
||||
minReplicas: {{ .Values.autoscaling.minReplicas }}
|
||||
maxReplicas: {{ .Values.autoscaling.maxReplicas }}
|
||||
{{- if .Values.autoscaling.behavior }}
|
||||
behavior:
|
||||
{{- toYaml .Values.autoscaling.behavior | nindent 4 }}
|
||||
{{- end }}
|
||||
metrics:
|
||||
{{- if .Values.autoscaling.targetCPUUtilizationPercentage }}
|
||||
- type: Resource
|
||||
|
|
|
|||
|
|
@ -41,7 +41,7 @@ spec:
|
|||
{{- end }}
|
||||
containers:
|
||||
- name: prisma-migrations
|
||||
image: "{{ .Values.image.repository }}:{{ .Values.image.tag | default (printf "main-%s" .Chart.AppVersion) }}"
|
||||
image: "{{ .Values.image.repository }}:{{ .Values.image.tag | default .Chart.AppVersion }}"
|
||||
imagePullPolicy: {{ .Values.image.pullPolicy }}
|
||||
securityContext:
|
||||
{{- toYaml .Values.securityContext | nindent 12 }}
|
||||
|
|
|
|||
36
deploy/charts/litellm-helm/tests/hpa_tests.yaml
Normal file
36
deploy/charts/litellm-helm/tests/hpa_tests.yaml
Normal file
|
|
@ -0,0 +1,36 @@
|
|||
suite: "hpa with behavior"
|
||||
templates:
|
||||
- hpa.yaml
|
||||
tests:
|
||||
- it: "renders behavior when set"
|
||||
set:
|
||||
autoscaling.enabled: true
|
||||
autoscaling.behavior:
|
||||
scaleUp:
|
||||
stabilizationWindowSeconds: 60
|
||||
policies:
|
||||
- type: Pods
|
||||
value: 2
|
||||
periodSeconds: 60
|
||||
scaleDown:
|
||||
stabilizationWindowSeconds: 90
|
||||
policies:
|
||||
- type: Pods
|
||||
value: 1
|
||||
periodSeconds: 60
|
||||
asserts:
|
||||
- isKind: { of: HorizontalPodAutoscaler }
|
||||
- equal: { path: spec.behavior.scaleUp.stabilizationWindowSeconds, value: 60 }
|
||||
- equal: { path: spec.behavior.scaleDown.stabilizationWindowSeconds, value: 90 }
|
||||
|
||||
---
|
||||
suite: "hpa without behavior"
|
||||
templates:
|
||||
- hpa.yaml
|
||||
tests:
|
||||
- it: "does not render behavior when not set"
|
||||
set:
|
||||
autoscaling.enabled: true
|
||||
asserts:
|
||||
- isKind: { of: HorizontalPodAutoscaler }
|
||||
- isNull: { path: spec.behavior }
|
||||
|
|
@ -10,7 +10,7 @@ image:
|
|||
repository: ghcr.io/berriai/litellm-database
|
||||
pullPolicy: Always
|
||||
# Overrides the image tag whose default is the chart appVersion.
|
||||
# tag: "main-latest"
|
||||
# tag: "latest"
|
||||
tag: ""
|
||||
|
||||
imagePullSecrets: []
|
||||
|
|
@ -184,6 +184,7 @@ autoscaling:
|
|||
maxReplicas: 100
|
||||
targetCPUUtilizationPercentage: 80
|
||||
# targetMemoryUtilizationPercentage: 80
|
||||
# behavior: {}
|
||||
|
||||
# Autoscaling with keda is mutually exclusive with hpa
|
||||
keda:
|
||||
|
|
|
|||
|
|
@ -24,7 +24,8 @@ RUN for i in 1 2 3; do \
|
|||
curl \
|
||||
openssl \
|
||||
libsndfile \
|
||||
nodejs && break || sleep 5; \
|
||||
nodejs \
|
||||
npm && break || sleep 5; \
|
||||
done
|
||||
|
||||
ENV UV_PROJECT_ENVIRONMENT=/app/.venv \
|
||||
|
|
|
|||
File diff suppressed because one or more lines are too long
|
|
@ -0,0 +1,4 @@
|
|||
-- AlterTable
|
||||
-- Adds the admin-toggleable pause flag used by the router's blocked filter and the
|
||||
-- credential lookup helpers; defaults to false so existing rows behave unchanged.
|
||||
ALTER TABLE "LiteLLM_ProxyModelTable" ADD COLUMN IF NOT EXISTS "blocked" BOOLEAN NOT NULL DEFAULT false;
|
||||
|
|
@ -48,9 +48,10 @@ model LiteLLM_CredentialsTable {
|
|||
// Models on proxy
|
||||
model LiteLLM_ProxyModelTable {
|
||||
model_id String @id @default(uuid())
|
||||
model_name String
|
||||
model_name String
|
||||
litellm_params Json
|
||||
model_info Json?
|
||||
model_info Json?
|
||||
blocked Boolean @default(false)
|
||||
created_at DateTime @default(now()) @map("created_at")
|
||||
created_by String
|
||||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.72"
|
||||
version = "0.4.73"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.4.72"
|
||||
version = "0.4.73"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -225,6 +225,10 @@ use_chat_completions_url_for_anthropic_messages: bool = bool(
|
|||
route_all_chat_openai_to_responses: bool = (
|
||||
os.getenv("LITELLM_ROUTE_ALL_CHAT_OPENAI_TO_RESPONSES", "false").lower() == "true"
|
||||
) # When True, routes all OpenAI /chat/completions requests through the Responses API bridge
|
||||
use_legacy_interactions_schema: bool = (
|
||||
os.getenv("LITELLM_USE_LEGACY_INTERACTIONS_SCHEMA", "false").lower() == "true"
|
||||
) # When True, sends Api-Revision: 2026-05-07 to Google so responses use the legacy `outputs`
|
||||
# schema instead of the new `steps` schema. Remove this flag after June 8, 2026.
|
||||
retry = True
|
||||
### AUTH ###
|
||||
api_key: Optional[str] = None
|
||||
|
|
@ -409,6 +413,12 @@ internal_user_budget_duration: Optional[str] = None
|
|||
tag_budget_config: Optional[Dict[str, "BudgetConfig"]] = None
|
||||
max_end_user_budget: Optional[float] = None
|
||||
max_end_user_budget_id: Optional[str] = None
|
||||
# When True, end-user IDs extracted from requests are validated against
|
||||
# LiteLLM_EndUserTable / LiteLLM_UserTable. Values that do not resolve to a
|
||||
# known row are dropped before reaching spend logs. Defaults to False for
|
||||
# backwards compatibility — arbitrary client-supplied identifiers still
|
||||
# pass through unchanged.
|
||||
validate_end_user_id_in_db: bool = False
|
||||
disable_end_user_cost_tracking: Optional[bool] = None
|
||||
disable_end_user_cost_tracking_prometheus_only: Optional[bool] = None
|
||||
enable_end_user_cost_tracking_prometheus_only: Optional[bool] = None
|
||||
|
|
@ -416,6 +426,7 @@ custom_prometheus_metadata_labels: List[str] = []
|
|||
custom_prometheus_tags: List[str] = []
|
||||
prometheus_metrics_config: Optional[List] = None
|
||||
prometheus_emit_stream_label: bool = False
|
||||
prometheus_user_budget_label_include_email_alias: bool = False
|
||||
prometheus_end_user_metrics_max_series_per_metric: Optional[int] = 10000
|
||||
prometheus_end_user_metrics_ttl_seconds: Optional[float] = 3600.0
|
||||
prometheus_end_user_metrics_cleanup_interval_seconds: Optional[float] = 60.0
|
||||
|
|
@ -631,6 +642,7 @@ minimax_models: Set = set()
|
|||
aws_polly_models: Set = set()
|
||||
gigachat_models: Set = set()
|
||||
llamagate_models: Set = set()
|
||||
reducto_models: Set = set()
|
||||
bedrock_mantle_models: Set = set()
|
||||
|
||||
|
||||
|
|
@ -898,6 +910,8 @@ def add_known_models(model_cost_map: Optional[Dict] = None):
|
|||
gigachat_models.add(key)
|
||||
elif value.get("litellm_provider") == "llamagate":
|
||||
llamagate_models.add(key)
|
||||
elif value.get("litellm_provider") == "reducto":
|
||||
reducto_models.add(key)
|
||||
elif value.get("litellm_provider") == "bedrock_mantle":
|
||||
bedrock_mantle_models.add(key)
|
||||
|
||||
|
|
@ -1009,6 +1023,7 @@ model_list = list(
|
|||
| ovhcloud_models
|
||||
| lemonade_models
|
||||
| docker_model_runner_models
|
||||
| reducto_models
|
||||
| bedrock_mantle_models
|
||||
| set(clarifai_models)
|
||||
)
|
||||
|
|
@ -1115,6 +1130,7 @@ models_by_provider: dict = {
|
|||
"aws_polly": aws_polly_models,
|
||||
"gigachat": gigachat_models,
|
||||
"llamagate": llamagate_models,
|
||||
"reducto": reducto_models,
|
||||
"bedrock_mantle": bedrock_mantle_models,
|
||||
}
|
||||
|
||||
|
|
@ -1287,6 +1303,18 @@ from .responses.main import *
|
|||
# Interactions API is available as litellm.interactions module
|
||||
# Usage: litellm.interactions.create(), litellm.interactions.get(), etc.
|
||||
from . import interactions
|
||||
from .interactions.agents.main import (
|
||||
acreate as acreate_agent,
|
||||
create as create_agent,
|
||||
alist as alist_agents,
|
||||
list as list_agents,
|
||||
aget as aget_agent,
|
||||
get as get_agent,
|
||||
adelete as adelete_agent,
|
||||
delete as delete_agent,
|
||||
alist_versions as alist_agent_versions,
|
||||
list_versions as list_agent_versions,
|
||||
)
|
||||
from .skills.main import (
|
||||
create_skill,
|
||||
acreate_skill,
|
||||
|
|
@ -1849,6 +1877,9 @@ if TYPE_CHECKING:
|
|||
from .llms.azure.completion.transformation import (
|
||||
AzureOpenAITextConfig as AzureOpenAITextConfig,
|
||||
)
|
||||
from .llms.azure.audio_transcription.transformation import (
|
||||
AzureSpeechAudioTranscriptionConfig as AzureSpeechAudioTranscriptionConfig,
|
||||
)
|
||||
from .llms.hosted_vllm.chat.transformation import (
|
||||
HostedVLLMChatConfig as HostedVLLMChatConfig,
|
||||
)
|
||||
|
|
@ -1880,6 +1911,12 @@ if TYPE_CHECKING:
|
|||
from .llms.dashscope.chat.transformation import (
|
||||
DashScopeChatConfig as DashScopeChatConfig,
|
||||
)
|
||||
from .llms.dashscope.embed.transformation import (
|
||||
DashScopeEmbeddingConfig as DashScopeEmbeddingConfig,
|
||||
)
|
||||
from .llms.dashscope.rerank.transformation import (
|
||||
DashScopeRerankConfig as DashScopeRerankConfig,
|
||||
)
|
||||
from .llms.moonshot.chat.transformation import (
|
||||
MoonshotChatConfig as MoonshotChatConfig,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -273,6 +273,7 @@ LLM_CONFIG_NAMES = (
|
|||
"AzureOpenAIConfig",
|
||||
"AzureOpenAIGPT5Config",
|
||||
"AzureOpenAITextConfig",
|
||||
"AzureSpeechAudioTranscriptionConfig",
|
||||
"HostedVLLMChatConfig",
|
||||
"HostedVLLMEmbeddingConfig",
|
||||
# Alias for backwards compatibility
|
||||
|
|
@ -1054,6 +1055,10 @@ _LLM_CONFIGS_IMPORT_MAP = {
|
|||
".llms.azure.completion.transformation",
|
||||
"AzureOpenAITextConfig",
|
||||
),
|
||||
"AzureSpeechAudioTranscriptionConfig": (
|
||||
".llms.azure.audio_transcription.transformation",
|
||||
"AzureSpeechAudioTranscriptionConfig",
|
||||
),
|
||||
"HostedVLLMChatConfig": (
|
||||
".llms.hosted_vllm.chat.transformation",
|
||||
"HostedVLLMChatConfig",
|
||||
|
|
|
|||
|
|
@ -100,6 +100,8 @@ def _get_redis_cluster_kwargs(client=None):
|
|||
"azure_tenant_id",
|
||||
"azure_client_secret",
|
||||
"max_connections",
|
||||
"socket_timeout",
|
||||
"socket_connect_timeout",
|
||||
}
|
||||
|
||||
return available_args
|
||||
|
|
|
|||
|
|
@ -6,7 +6,6 @@ Always uses fastuuid for performance.
|
|||
|
||||
import fastuuid as _uuid # type: ignore
|
||||
|
||||
|
||||
# Expose a module-like alias so callers can use: uuid.uuid4()
|
||||
uuid = _uuid
|
||||
|
||||
|
|
|
|||
|
|
@ -9,7 +9,6 @@ from typing import Dict, Optional
|
|||
|
||||
from .exceptions import AnthropicErrorResponse, AnthropicErrorType
|
||||
|
||||
|
||||
# HTTP status code -> Anthropic error type
|
||||
# Source: https://docs.anthropic.com/en/api/errors
|
||||
ANTHROPIC_ERROR_TYPE_MAP: Dict[int, AnthropicErrorType] = {
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@
|
|||
|
||||
from typing_extensions import Literal, Required, TypedDict
|
||||
|
||||
|
||||
# Known Anthropic error types
|
||||
# Source: https://docs.anthropic.com/en/api/errors
|
||||
AnthropicErrorType = Literal[
|
||||
|
|
|
|||
|
|
@ -87,6 +87,16 @@ class CachingHandlerResponse(BaseModel):
|
|||
in_memory_cache_obj = InMemoryCache()
|
||||
|
||||
|
||||
def _is_chat_completion_cached_dict(cached_result: dict) -> bool:
|
||||
cached_id = cached_result.get("id")
|
||||
if isinstance(cached_id, str) and cached_id.startswith("chatcmpl"):
|
||||
return True
|
||||
obj = cached_result.get("object")
|
||||
if isinstance(obj, str):
|
||||
return obj.startswith("chat.completion")
|
||||
return "choices" in cached_result
|
||||
|
||||
|
||||
def _should_defer_streaming_cache_hit_callbacks(*, kwargs: Dict[str, Any]) -> bool:
|
||||
"""
|
||||
When stream=True, do not run success callbacks at cache-hit time.
|
||||
|
|
@ -861,27 +871,47 @@ class LLMCachingHandler:
|
|||
elif (call_type == "aresponses" or call_type == "responses") and isinstance(
|
||||
cached_result, dict
|
||||
):
|
||||
from litellm.responses.streaming_iterator import (
|
||||
CachedResponsesAPIStreamingIterator,
|
||||
)
|
||||
|
||||
response_obj = ResponsesAPIResponse(**cached_result)
|
||||
if (
|
||||
hasattr(response_obj, "_hidden_params")
|
||||
and response_obj._hidden_params is not None
|
||||
and isinstance(response_obj._hidden_params, dict)
|
||||
):
|
||||
response_obj._hidden_params["cache_hit"] = True
|
||||
|
||||
if kwargs.get("stream", False) is True:
|
||||
cached_result = CachedResponsesAPIStreamingIterator(
|
||||
response=response_obj,
|
||||
logging_obj=logging_obj,
|
||||
request_data=kwargs,
|
||||
call_type=call_type,
|
||||
)
|
||||
use_chat_completion_cache = _is_chat_completion_cached_dict(cached_result)
|
||||
if use_chat_completion_cache:
|
||||
if kwargs.get("stream", False) is True:
|
||||
bridge_call_type = (
|
||||
CallTypes.acompletion.value
|
||||
if call_type == "aresponses"
|
||||
else CallTypes.completion.value
|
||||
)
|
||||
cached_result = self._convert_cached_stream_response(
|
||||
cached_result=cached_result,
|
||||
call_type=bridge_call_type,
|
||||
logging_obj=logging_obj,
|
||||
model=model,
|
||||
)
|
||||
else:
|
||||
cached_result = convert_to_model_response_object(
|
||||
response_object=cached_result,
|
||||
model_response_object=ModelResponse(),
|
||||
)
|
||||
else:
|
||||
cached_result = response_obj
|
||||
from litellm.responses.streaming_iterator import (
|
||||
CachedResponsesAPIStreamingIterator,
|
||||
)
|
||||
|
||||
response_obj = ResponsesAPIResponse(**cached_result)
|
||||
if (
|
||||
hasattr(response_obj, "_hidden_params")
|
||||
and response_obj._hidden_params is not None
|
||||
and isinstance(response_obj._hidden_params, dict)
|
||||
):
|
||||
response_obj._hidden_params["cache_hit"] = True
|
||||
|
||||
if kwargs.get("stream", False) is True:
|
||||
cached_result = CachedResponsesAPIStreamingIterator(
|
||||
response=response_obj,
|
||||
logging_obj=logging_obj,
|
||||
request_data=kwargs,
|
||||
call_type=call_type,
|
||||
)
|
||||
else:
|
||||
cached_result = response_obj
|
||||
|
||||
if (
|
||||
hasattr(cached_result, "_hidden_params")
|
||||
|
|
|
|||
|
|
@ -37,6 +37,15 @@ class ResponsesToCompletionBridgeHandler:
|
|||
stream = litellm_params.get("stream", False)
|
||||
return bool(stream)
|
||||
|
||||
@staticmethod
|
||||
def _is_preformatted_cached_chat_stream(result: Any) -> bool:
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
|
||||
return (
|
||||
isinstance(result, CustomStreamWrapper)
|
||||
and result.custom_llm_provider == "cached_response"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _coerce_response_object(
|
||||
response_obj: Any,
|
||||
|
|
@ -177,6 +186,8 @@ class ResponsesToCompletionBridgeHandler:
|
|||
**request_data,
|
||||
)
|
||||
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
stream = self._resolve_stream_flag(optional_params, litellm_params)
|
||||
if isinstance(result, ResponsesAPIResponse):
|
||||
return self.transformation_handler.transform_response(
|
||||
|
|
@ -192,6 +203,8 @@ class ResponsesToCompletionBridgeHandler:
|
|||
api_key=kwargs.get("api_key"),
|
||||
json_mode=kwargs.get("json_mode"),
|
||||
)
|
||||
elif isinstance(result, ModelResponse):
|
||||
return result
|
||||
elif not stream:
|
||||
responses_api_response = self._collect_response_from_stream(result)
|
||||
return self.transformation_handler.transform_response(
|
||||
|
|
@ -208,6 +221,10 @@ class ResponsesToCompletionBridgeHandler:
|
|||
json_mode=kwargs.get("json_mode"),
|
||||
)
|
||||
else:
|
||||
if self._is_preformatted_cached_chat_stream(result):
|
||||
return self._apply_post_stream_processing(
|
||||
result, model, custom_llm_provider
|
||||
)
|
||||
completion_stream = self.transformation_handler.get_model_response_iterator(
|
||||
streaming_response=result, # type: ignore
|
||||
sync_stream=True,
|
||||
|
|
@ -256,6 +273,8 @@ class ResponsesToCompletionBridgeHandler:
|
|||
aresponses=True,
|
||||
)
|
||||
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
stream = self._resolve_stream_flag(optional_params, litellm_params)
|
||||
if isinstance(result, ResponsesAPIResponse):
|
||||
return self.transformation_handler.transform_response(
|
||||
|
|
@ -271,6 +290,8 @@ class ResponsesToCompletionBridgeHandler:
|
|||
api_key=kwargs.get("api_key"),
|
||||
json_mode=kwargs.get("json_mode"),
|
||||
)
|
||||
elif isinstance(result, ModelResponse):
|
||||
return result
|
||||
elif not stream:
|
||||
responses_api_response = await self._collect_response_from_stream_async(
|
||||
result
|
||||
|
|
@ -289,6 +310,10 @@ class ResponsesToCompletionBridgeHandler:
|
|||
json_mode=kwargs.get("json_mode"),
|
||||
)
|
||||
else:
|
||||
if self._is_preformatted_cached_chat_stream(result):
|
||||
return self._apply_post_stream_processing(
|
||||
result, model, custom_llm_provider
|
||||
)
|
||||
completion_stream = self.transformation_handler.get_model_response_iterator(
|
||||
streaming_response=result, # type: ignore
|
||||
sync_stream=False,
|
||||
|
|
|
|||
|
|
@ -30,6 +30,11 @@ from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
|||
from litellm.llms.base_llm.bridges.completion_transformation import (
|
||||
CompletionTransformationBridge,
|
||||
)
|
||||
from litellm.responses.sse_output_recovery import (
|
||||
parse_sse_json_chunk,
|
||||
record_output_item_chunk,
|
||||
record_output_text_chunk,
|
||||
)
|
||||
from litellm.types.llms.openai import (
|
||||
ChatCompletionAnnotation,
|
||||
ChatCompletionReasoningItem,
|
||||
|
|
@ -97,7 +102,7 @@ def _build_reasoning_item(
|
|||
|
||||
|
||||
def _reasoning_item_to_response_input(
|
||||
r_item: Union[ChatCompletionReasoningItem, Dict[str, Any]]
|
||||
r_item: Union[ChatCompletionReasoningItem, Dict[str, Any]],
|
||||
) -> Dict[str, Any]:
|
||||
"""Convert a stored ChatCompletionReasoningItem back to a Responses API input item."""
|
||||
r_input: Dict[str, Any] = {
|
||||
|
|
@ -601,6 +606,79 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
|
||||
return choices
|
||||
|
||||
@classmethod
|
||||
def _extract_output_from_completed_event(
|
||||
cls, parsed_chunk: Dict[str, Any]
|
||||
) -> Optional[List[Dict[str, Any]]]:
|
||||
response_payload = parsed_chunk.get("response")
|
||||
if not isinstance(response_payload, dict):
|
||||
return None
|
||||
response_output = response_payload.get("output")
|
||||
if not isinstance(response_output, list) or len(response_output) == 0:
|
||||
return None
|
||||
return cast(List[Dict[str, Any]], response_output)
|
||||
|
||||
@classmethod
|
||||
def _recover_output_items_from_raw_sse(
|
||||
cls, raw_sse: Optional[str]
|
||||
) -> List[Dict[str, Any]]:
|
||||
if not raw_sse or not isinstance(raw_sse, str):
|
||||
return []
|
||||
|
||||
recovered_output_items: Dict[int, Dict[str, Any]] = {}
|
||||
recovered_text_only_items: Dict[int, Dict[str, Any]] = {}
|
||||
|
||||
for chunk in raw_sse.splitlines():
|
||||
parsed_chunk = parse_sse_json_chunk(chunk)
|
||||
if parsed_chunk is None:
|
||||
continue
|
||||
|
||||
event_type = parsed_chunk.get("type")
|
||||
|
||||
if event_type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED:
|
||||
recovered_output = cls._extract_output_from_completed_event(
|
||||
parsed_chunk
|
||||
)
|
||||
if recovered_output is not None:
|
||||
return recovered_output
|
||||
continue
|
||||
|
||||
if event_type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE:
|
||||
record_output_item_chunk(
|
||||
parsed_chunk=parsed_chunk,
|
||||
output_items=recovered_output_items,
|
||||
)
|
||||
continue
|
||||
|
||||
if event_type == ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE:
|
||||
record_output_text_chunk(
|
||||
parsed_chunk=parsed_chunk,
|
||||
output_items=recovered_output_items,
|
||||
text_only_items=recovered_text_only_items,
|
||||
)
|
||||
continue
|
||||
|
||||
# Merge text-only items into the recovered output items. Real
|
||||
# OUTPUT_ITEM_DONE events take precedence at any given output_index,
|
||||
# but text-only items at indices without a matching OUTPUT_ITEM_DONE
|
||||
# must still be preserved (e.g. multi-output responses where some
|
||||
# indices only emitted OUTPUT_TEXT_DONE).
|
||||
merged_items: Dict[int, Dict[str, Any]] = {**recovered_text_only_items}
|
||||
merged_items.update(recovered_output_items)
|
||||
|
||||
if merged_items:
|
||||
return [item for _, item in sorted(merged_items.items())]
|
||||
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
def _recover_output_items_from_logging(
|
||||
cls, logging_obj: "LiteLLMLoggingObj"
|
||||
) -> List[Dict[str, Any]]:
|
||||
model_call_details = getattr(logging_obj, "model_call_details", {}) or {}
|
||||
original_response = model_call_details.get("original_response")
|
||||
return cls._recover_output_items_from_raw_sse(original_response)
|
||||
|
||||
def transform_response( # noqa: PLR0915
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -625,9 +703,22 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
if raw_response.error is not None:
|
||||
raise ValueError(f"Error in response: {raw_response.error}")
|
||||
|
||||
output_items = raw_response.output
|
||||
if len(output_items) == 0:
|
||||
recovered_output_items = self._recover_output_items_from_logging(
|
||||
logging_obj
|
||||
)
|
||||
if recovered_output_items:
|
||||
output_items = cast(Any, recovered_output_items)
|
||||
raw_response.output = cast(Any, recovered_output_items)
|
||||
verbose_logger.warning(
|
||||
"Recovered empty Responses API output from raw SSE for model=%s",
|
||||
model,
|
||||
)
|
||||
|
||||
# Convert response output to choices using the static helper
|
||||
choices = self._convert_response_output_to_choices(
|
||||
output_items=raw_response.output,
|
||||
output_items=output_items,
|
||||
handle_raw_dict_callback=self._handle_raw_dict_response_item,
|
||||
)
|
||||
|
||||
|
|
@ -641,7 +732,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unknown items in responses API response: {raw_response.output}"
|
||||
f"Unknown items in responses API response: {output_items}"
|
||||
)
|
||||
|
||||
setattr(model_response, "choices", choices)
|
||||
|
|
@ -1141,6 +1232,14 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
event_type = parsed_chunk.get("type")
|
||||
if isinstance(event_type, ResponsesAPIStreamEvents):
|
||||
event_type = event_type.value
|
||||
|
||||
if parsed_chunk.get("object") == "chat.completion.chunk" or (
|
||||
event_type is None
|
||||
and isinstance(parsed_chunk.get("choices"), list)
|
||||
and parsed_chunk.get("choices")
|
||||
):
|
||||
return ModelResponseStream(**parsed_chunk)
|
||||
|
||||
verbose_logger.debug(f"Chat provider: Processing event type: {event_type}")
|
||||
|
||||
if event_type == "response.created":
|
||||
|
|
@ -1229,7 +1328,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
raise ValueError(
|
||||
f"Chat provider: Invalid function argument delta {parsed_chunk}"
|
||||
)
|
||||
elif event_type == "response.output_item.done":
|
||||
elif event_type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE:
|
||||
# New output item added
|
||||
output_item = parsed_chunk.get("item", {})
|
||||
if output_item.get("type") == "function_call":
|
||||
|
|
|
|||
|
|
@ -5,7 +5,6 @@ Auto-detect content type per message: code, JSON, or text.
|
|||
import json
|
||||
import re
|
||||
|
||||
|
||||
_CODE_KEYWORDS = re.compile(
|
||||
r"\b(?:def |function |class |import |from |require\(|#include|fn |func |const |let |var |public |private |static )\b"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1443,6 +1443,12 @@ CLI_JWT_EXPIRATION_HOURS = int(
|
|||
or os.getenv("LITELLM_CLI_JWT_EXPIRATION_HOURS")
|
||||
or 24
|
||||
)
|
||||
# Comma-separated allowlisted OIDC claim map for CLI SSO polling, e.g.
|
||||
# "employment_type->acme_employment_type,org_info.department->department"
|
||||
CLI_SSO_CLAIM_MAP = (
|
||||
os.getenv("CLI_SSO_CLAIM_MAP") or os.getenv("LITELLM_CLI_SSO_CLAIM_MAP") or ""
|
||||
)
|
||||
CLI_SSO_CLAIM_MAX_SCALAR_LENGTH = 1024
|
||||
|
||||
########################### UI SESSION DURATION ###########################
|
||||
# Duration for UI login session (username/password, SSO, invitation links). Format: "30s", "30m", "24h", "7d"
|
||||
|
|
|
|||
|
|
@ -173,17 +173,45 @@ def _cost_per_token_custom_pricing_helper(
|
|||
prompt_tokens: float = 0,
|
||||
completion_tokens: float = 0,
|
||||
response_time_ms: Optional[float] = 0.0,
|
||||
cached_tokens: float = 0,
|
||||
cache_creation_tokens: float = 0,
|
||||
### CUSTOM PRICING ###
|
||||
custom_cost_per_token: Optional[CostPerToken] = None,
|
||||
custom_cost_per_second: Optional[float] = None,
|
||||
) -> Optional[Tuple[float, float]]:
|
||||
"""Internal helper function for calculating cost, if custom pricing given"""
|
||||
"""Internal helper function for calculating cost, if custom pricing given.
|
||||
|
||||
prompt_tokens is assumed to include both cached_tokens and cache_creation_tokens
|
||||
(OpenAI-compatible convention). Anthropic-style usage where prompt_tokens excludes
|
||||
cache tokens is handled at the caller (cost_per_token) before invoking this helper.
|
||||
"""
|
||||
if custom_cost_per_token is None and custom_cost_per_second is None:
|
||||
return None
|
||||
|
||||
if custom_cost_per_token is not None:
|
||||
input_cost = custom_cost_per_token["input_cost_per_token"] * prompt_tokens
|
||||
output_cost = custom_cost_per_token["output_cost_per_token"] * completion_tokens
|
||||
input_cost_per_token = custom_cost_per_token["input_cost_per_token"]
|
||||
output_cost_per_token = custom_cost_per_token["output_cost_per_token"]
|
||||
|
||||
cache_read_input_token_cost = custom_cost_per_token.get(
|
||||
"cache_read_input_token_cost",
|
||||
input_cost_per_token,
|
||||
)
|
||||
cache_creation_input_token_cost = custom_cost_per_token.get(
|
||||
"cache_creation_input_token_cost",
|
||||
input_cost_per_token,
|
||||
)
|
||||
|
||||
regular_prompt_tokens = max(
|
||||
prompt_tokens - cached_tokens - cache_creation_tokens,
|
||||
0,
|
||||
)
|
||||
|
||||
input_cost = (
|
||||
regular_prompt_tokens * input_cost_per_token
|
||||
+ cached_tokens * cache_read_input_token_cost
|
||||
+ cache_creation_tokens * cache_creation_input_token_cost
|
||||
)
|
||||
output_cost = completion_tokens * output_cost_per_token
|
||||
return input_cost, output_cost
|
||||
elif custom_cost_per_second is not None:
|
||||
output_cost = custom_cost_per_second * response_time_ms / 1000 # type: ignore
|
||||
|
|
@ -323,10 +351,56 @@ def cost_per_token( # noqa: PLR0915
|
|||
)
|
||||
|
||||
## CUSTOM PRICING ##
|
||||
# Normalize cache token counts across providers:
|
||||
# - OpenAI-compatible: usage.prompt_tokens_details.cached_tokens
|
||||
# (prompt_tokens already INCLUDES cached_tokens)
|
||||
# - Anthropic: usage.cache_read_input_tokens / cache_creation_input_tokens
|
||||
# (prompt_tokens does NOT include these — adjust before calling helper)
|
||||
_cache_read_tokens: float = 0
|
||||
_cache_creation_tokens: float = 0
|
||||
_is_anthropic_style = False
|
||||
|
||||
if usage_object is not None:
|
||||
_pt_details = getattr(usage_object, "prompt_tokens_details", None)
|
||||
if _pt_details is not None:
|
||||
_cache_read_tokens = float(getattr(_pt_details, "cached_tokens", 0) or 0)
|
||||
# OpenAI-compatible providers report cache-write tokens under
|
||||
# either `cache_write_tokens` (kimi-k2) or `cache_creation_tokens`.
|
||||
# Mirror db_spend_update_writer to stay symmetric.
|
||||
_cache_creation_tokens = float(
|
||||
getattr(_pt_details, "cache_write_tokens", 0)
|
||||
or getattr(_pt_details, "cache_creation_tokens", 0)
|
||||
or 0
|
||||
)
|
||||
|
||||
_anthropic_read = getattr(usage_object, "cache_read_input_tokens", None)
|
||||
_anthropic_create = getattr(usage_object, "cache_creation_input_tokens", None)
|
||||
if _anthropic_read is not None or _anthropic_create is not None:
|
||||
_is_anthropic_style = True
|
||||
if _anthropic_read is not None:
|
||||
_cache_read_tokens = float(_anthropic_read)
|
||||
if _anthropic_create is not None:
|
||||
_cache_creation_tokens = float(_anthropic_create)
|
||||
|
||||
if not _cache_read_tokens and cache_read_input_tokens:
|
||||
_cache_read_tokens = float(cache_read_input_tokens)
|
||||
_is_anthropic_style = True
|
||||
if not _cache_creation_tokens and cache_creation_input_tokens:
|
||||
_cache_creation_tokens = float(cache_creation_input_tokens)
|
||||
_is_anthropic_style = True
|
||||
|
||||
# Anthropic reports prompt_tokens as input_tokens (excluding cache tokens).
|
||||
# Adjust so the helper's "prompt_tokens includes cache tokens" invariant holds.
|
||||
_normalized_prompt_tokens = float(prompt_tokens)
|
||||
if _is_anthropic_style:
|
||||
_normalized_prompt_tokens += _cache_read_tokens + _cache_creation_tokens
|
||||
|
||||
response_cost = _cost_per_token_custom_pricing_helper(
|
||||
prompt_tokens=prompt_tokens,
|
||||
prompt_tokens=_normalized_prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
response_time_ms=response_time_ms,
|
||||
cached_tokens=_cache_read_tokens,
|
||||
cache_creation_tokens=_cache_creation_tokens,
|
||||
custom_cost_per_second=custom_cost_per_second,
|
||||
custom_cost_per_token=custom_cost_per_token,
|
||||
)
|
||||
|
|
@ -1805,10 +1879,6 @@ def ocr_cost(
|
|||
if response.usage_info is None:
|
||||
raise ValueError("OCR response usage_info is None")
|
||||
|
||||
pages_processed = response.usage_info.pages_processed
|
||||
if pages_processed is None:
|
||||
raise ValueError("OCR response pages_processed is None")
|
||||
|
||||
try:
|
||||
model_info: Optional[ModelInfo] = litellm.get_model_info(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
|
|
@ -1816,9 +1886,49 @@ def ocr_cost(
|
|||
except Exception:
|
||||
model_info = None
|
||||
|
||||
ocr_cost_per_page: float = 0.0
|
||||
credits = getattr(response.usage_info, "credits", None)
|
||||
cost_per_credit = None
|
||||
if model_info is not None:
|
||||
ocr_cost_per_page = model_info.get("ocr_cost_per_page") or 0.0
|
||||
cost_per_credit = model_info.get("ocr_cost_per_credit")
|
||||
if credits is not None and cost_per_credit is not None:
|
||||
return cost_per_credit * credits, 0.0
|
||||
|
||||
ocr_cost_per_page: Optional[float] = None
|
||||
if model_info is not None:
|
||||
ocr_cost_per_page = model_info.get("ocr_cost_per_page")
|
||||
|
||||
pages_processed = response.usage_info.pages_processed
|
||||
if pages_processed is None:
|
||||
if cost_per_credit is not None or ocr_cost_per_page is None:
|
||||
# Surface missing usage data instead of silently under-reporting
|
||||
# cost. The previous behavior raised ValueError; we now return 0.0
|
||||
# for credit-priced or unpriced models, so log a warning to keep
|
||||
# the regression visible to operators.
|
||||
verbose_logger.warning(
|
||||
"OCR cost: model=%s custom_llm_provider=%s response.usage_info."
|
||||
"pages_processed is None and credits=%s; returning 0.0 cost.",
|
||||
model,
|
||||
custom_llm_provider,
|
||||
credits,
|
||||
)
|
||||
return 0.0, 0.0
|
||||
raise ValueError("OCR response pages_processed is None")
|
||||
|
||||
if ocr_cost_per_page is None:
|
||||
# No per-page pricing configured. Either the model is on credit-based
|
||||
# pricing (and credits weren't returned, so the credit branch above did
|
||||
# not match) or the model has no OCR pricing entry at all. Surface a
|
||||
# warning so that missing pricing entries are visible rather than
|
||||
# silently producing zero cost for billable usage.
|
||||
verbose_logger.warning(
|
||||
"OCR cost: model=%s custom_llm_provider=%s reported "
|
||||
"pages_processed=%s but no ocr_cost_per_page is configured; "
|
||||
"returning 0.0 cost.",
|
||||
model,
|
||||
custom_llm_provider,
|
||||
pages_processed,
|
||||
)
|
||||
return 0.0, 0.0
|
||||
|
||||
total_ocr_processing_cost: float = ocr_cost_per_page * pages_processed
|
||||
return total_ocr_processing_cost, 0.0
|
||||
|
|
|
|||
|
|
@ -918,9 +918,11 @@ class GuardrailRaisedException(Exception):
|
|||
guardrail_name: Optional[str] = None,
|
||||
message: str = "",
|
||||
should_wrap_with_default_message: bool = True,
|
||||
status_code: int = 400,
|
||||
):
|
||||
default_message = f"Guardrail raised an exception, Guardrail: {guardrail_name}, Message: {message}"
|
||||
self.guardrail_name = guardrail_name
|
||||
self.status_code = status_code
|
||||
self.message = default_message if should_wrap_with_default_message else message
|
||||
super().__init__(self.message)
|
||||
|
||||
|
|
@ -930,12 +932,14 @@ class BlockedPiiEntityError(Exception):
|
|||
self,
|
||||
entity_type: str,
|
||||
guardrail_name: Optional[str] = None,
|
||||
status_code: int = 400,
|
||||
):
|
||||
"""
|
||||
Raised when a blocked entity is detected by a guardrail.
|
||||
"""
|
||||
self.entity_type = entity_type
|
||||
self.guardrail_name = guardrail_name
|
||||
self.status_code = status_code
|
||||
self.message = f"Blocked entity detected: {entity_type} by Guardrail: {guardrail_name}. This entity is not allowed to be used in this request."
|
||||
super().__init__(self.message)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
from typing import AsyncIterator, Dict, Iterator, Literal, NamedTuple, Union
|
||||
|
||||
|
||||
FileContentProvider = Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic", "manus"
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,10 +1,10 @@
|
|||
"""
|
||||
Google GenAI Adapters for LiteLLM
|
||||
|
||||
This module provides adapters for transforming Google GenAI generate_content requests
|
||||
This module provides adapters for transforming Google GenAI generate_content requests
|
||||
to/from LiteLLM completion format with full support for:
|
||||
- Text content transformation
|
||||
- Tool calling (function declarations, function calls, function responses)
|
||||
- Tool calling (function declarations, function calls, function responses)
|
||||
- Streaming (both regular and tool calling)
|
||||
- Mixed content (text + tool calls)
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
"""
|
||||
Handles Batching + sending Httpx Post requests to slack
|
||||
Handles Batching + sending Httpx Post requests to slack
|
||||
|
||||
Slack alerts are sent every 10s or when events are greater than X events
|
||||
Slack alerts are sent every 10s or when events are greater than X events
|
||||
|
||||
see custom_batch_logger.py for more details / defaults
|
||||
see custom_batch_logger.py for more details / defaults
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ else:
|
|||
|
||||
|
||||
def process_slack_alerting_variables(
|
||||
alert_to_webhook_url: Optional[Dict[AlertType, Union[List[str], str]]]
|
||||
alert_to_webhook_url: Optional[Dict[AlertType, Union[List[str], str]]],
|
||||
) -> Optional[Dict[AlertType, Union[List[str], str]]]:
|
||||
"""
|
||||
process alert_to_webhook_url
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Base class for Additional Logging Utils for CustomLoggers
|
||||
Base class for Additional Logging Utils for CustomLoggers
|
||||
|
||||
- Health Check for the logging util
|
||||
- Get Request / Response Payload for the logging util
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Custom Logger that handles batching logic
|
||||
Custom Logger that handles batching logic
|
||||
|
||||
Use this if you want your logs to be stored in memory and flushed periodically.
|
||||
"""
|
||||
|
|
@ -14,22 +14,38 @@ from litellm.integrations.custom_logger import CustomLogger
|
|||
|
||||
|
||||
class CustomBatchLogger(CustomLogger):
|
||||
preserve_events_added_during_flush = False
|
||||
|
||||
# Default cap on the in-memory log queue. Prevents unbounded memory growth
|
||||
# if ``async_send_batch`` consistently fails (e.g. the destination is
|
||||
# unreachable) and events are preserved across flush attempts. Subclasses
|
||||
# may override by passing ``max_queue_size`` or by setting the attribute
|
||||
# directly (see ``RubrikLogger`` for an example).
|
||||
DEFAULT_MAX_QUEUE_SIZE = 50_000
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
flush_lock: Optional[asyncio.Lock] = None,
|
||||
batch_size: Optional[int] = None,
|
||||
flush_interval: Optional[int] = None,
|
||||
max_queue_size: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
"""
|
||||
Args:
|
||||
flush_lock (Optional[asyncio.Lock], optional): Lock to use when flushing the queue. Defaults to None. Only used for custom loggers that do batching
|
||||
max_queue_size (Optional[int], optional): Maximum number of events to retain in ``log_queue``. When the limit is exceeded (e.g. because the send destination is unreachable and events are preserved for retry), the oldest events are dropped. Defaults to ``DEFAULT_MAX_QUEUE_SIZE``.
|
||||
"""
|
||||
self.log_queue: List = []
|
||||
self.flush_interval = flush_interval or litellm.DEFAULT_FLUSH_INTERVAL_SECONDS
|
||||
self.batch_size: int = batch_size or litellm.DEFAULT_BATCH_SIZE
|
||||
self.last_flush_time = time.time()
|
||||
self.flush_lock = flush_lock
|
||||
self.max_queue_size: int = (
|
||||
max_queue_size
|
||||
if max_queue_size is not None
|
||||
else self.DEFAULT_MAX_QUEUE_SIZE
|
||||
)
|
||||
|
||||
super().__init__(**kwargs)
|
||||
|
||||
|
|
@ -47,11 +63,40 @@ class CustomBatchLogger(CustomLogger):
|
|||
|
||||
async with self.flush_lock:
|
||||
if self.log_queue:
|
||||
log_queue_length = len(self.log_queue)
|
||||
verbose_logger.debug(
|
||||
"CustomLogger: Flushing batch of %s events", len(self.log_queue)
|
||||
)
|
||||
await self.async_send_batch()
|
||||
self.log_queue.clear()
|
||||
try:
|
||||
await self.async_send_batch()
|
||||
except Exception:
|
||||
# If the underlying batch send raised, do NOT drop the
|
||||
# in-flight events. They will be retried on the next flush.
|
||||
# Most existing async_send_batch implementations swallow
|
||||
# their own errors, so this only affects loggers that opt
|
||||
# in to surfacing failures (e.g. Rubrik).
|
||||
verbose_logger.exception(
|
||||
"CustomLogger: async_send_batch raised; preserving "
|
||||
"%s events in queue for retry",
|
||||
log_queue_length,
|
||||
)
|
||||
# Guard against unbounded queue growth if the destination
|
||||
# is persistently unreachable. Drop the oldest events
|
||||
# beyond ``max_queue_size``.
|
||||
overflow = len(self.log_queue) - self.max_queue_size
|
||||
if overflow > 0:
|
||||
del self.log_queue[:overflow]
|
||||
verbose_logger.warning(
|
||||
"CustomLogger: log queue exceeded max_queue_size=%s; "
|
||||
"dropped %s oldest events.",
|
||||
self.max_queue_size,
|
||||
overflow,
|
||||
)
|
||||
return
|
||||
if self.preserve_events_added_during_flush:
|
||||
del self.log_queue[:log_queue_length]
|
||||
else:
|
||||
self.log_queue.clear()
|
||||
self.last_flush_time = time.time()
|
||||
|
||||
async def async_send_batch(self, *args, **kwargs):
|
||||
|
|
|
|||
|
|
@ -43,7 +43,11 @@ if TYPE_CHECKING:
|
|||
dc = DualCache()
|
||||
|
||||
|
||||
from litellm.exceptions import ModifyResponseException as ModifyResponseException
|
||||
from litellm.exceptions import (
|
||||
BlockedPiiEntityError,
|
||||
GuardrailRaisedException,
|
||||
ModifyResponseException,
|
||||
)
|
||||
|
||||
|
||||
class CustomGuardrail(CustomLogger):
|
||||
|
|
@ -737,12 +741,15 @@ class CustomGuardrail(CustomLogger):
|
|||
(this was logged previously as an API failure - guardrail_failed_to_respond).
|
||||
|
||||
Guardrails signal intentional blocks by raising:
|
||||
- GuardrailRaisedException (generic guardrail API, tool permission)
|
||||
- BlockedPiiEntityError (Presidio PII detection)
|
||||
- HTTPException with status 400 (content policy violation)
|
||||
- ModifyResponseException (passthrough mode violation)
|
||||
"""
|
||||
|
||||
if isinstance(e, ModifyResponseException):
|
||||
return True
|
||||
if isinstance(e, (GuardrailRaisedException, BlockedPiiEntityError)):
|
||||
return True
|
||||
if (
|
||||
HTTPException is not None
|
||||
and isinstance(e, HTTPException)
|
||||
|
|
|
|||
|
|
@ -9,7 +9,6 @@ import polars as pl
|
|||
|
||||
from .schema import FOCUS_NORMALIZED_SCHEMA
|
||||
|
||||
|
||||
_TAG_KEYS = (
|
||||
"team_id",
|
||||
"team_alias",
|
||||
|
|
|
|||
|
|
@ -673,6 +673,15 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
if parent_otel_span is not None:
|
||||
parent_otel_span.set_status(Status(StatusCode.ERROR))
|
||||
|
||||
# Stamp team attributes onto the SERVER (root) span too, so the
|
||||
# trace root is team-filterable on the failure path like the
|
||||
# child exception span below.
|
||||
self._set_team_attributes_on_span(
|
||||
span=parent_otel_span,
|
||||
team_id=user_api_key_dict.team_id,
|
||||
team_alias=user_api_key_dict.team_alias,
|
||||
)
|
||||
|
||||
# Stamp structured error attrs on the SERVER span itself; the
|
||||
# failure path otherwise only sets its status (_handle_failure
|
||||
# records on the litellm_request child span). Inline import:
|
||||
|
|
@ -693,6 +702,14 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
},
|
||||
)
|
||||
|
||||
# _record_exception_on_span only stamps when error_code is set;
|
||||
# bare TypeError etc. has none, and the span is about to be ended.
|
||||
error_code = (
|
||||
error_information.get("error_code") if error_information else None
|
||||
)
|
||||
if not error_code:
|
||||
self.set_response_status_code_attribute(parent_otel_span, 500)
|
||||
|
||||
# Pre-request latency (request_data carries the propagated
|
||||
# metadata on the failure path; omitted if it failed before handoff).
|
||||
self.set_preprocessing_duration_attribute(parent_otel_span, request_data)
|
||||
|
|
@ -709,12 +726,65 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
key="exception",
|
||||
value=str(original_exception),
|
||||
)
|
||||
self._set_team_attributes_on_span(
|
||||
span=exception_logging_span,
|
||||
team_id=user_api_key_dict.team_id,
|
||||
team_alias=user_api_key_dict.team_alias,
|
||||
)
|
||||
exception_logging_span.set_status(Status(StatusCode.ERROR))
|
||||
exception_logging_span.end(end_time=self._to_ns(datetime.now()))
|
||||
|
||||
# Emit guardrail spans for any guardrail invocations that
|
||||
# ran during this request. _handle_failure typically does this,
|
||||
# but for pre-call guardrail blocks the standard_logging_object
|
||||
# may not carry guardrail_information by the time _handle_failure
|
||||
# fires (the data lives only in request_data["metadata"]). Pull
|
||||
# directly from request_data so the span is recorded either way;
|
||||
# _emit_once dedupes if _handle_failure already emitted it.
|
||||
self._emit_guardrail_spans_from_request_data(
|
||||
request_data=request_data,
|
||||
parent_span=parent_otel_span,
|
||||
)
|
||||
|
||||
# End Parent OTEL Sspan
|
||||
parent_otel_span.end(end_time=self._to_ns(datetime.now()))
|
||||
|
||||
def _emit_guardrail_spans_from_request_data(
|
||||
self,
|
||||
request_data: dict,
|
||||
parent_span: Optional[Any],
|
||||
) -> None:
|
||||
"""Emit ``guardrail`` spans from ``request_data["metadata"]
|
||||
["standard_logging_guardrail_information"]``.
|
||||
|
||||
Routed through ``_create_guardrail_span`` so the dedupe state in
|
||||
``_otel_internal`` is honoured — if ``_handle_failure`` already
|
||||
emitted these spans for the same kwargs, this is a no-op.
|
||||
"""
|
||||
from opentelemetry import trace as _trace
|
||||
|
||||
metadata = (request_data or {}).get("metadata") or {}
|
||||
guardrail_information = metadata.get("standard_logging_guardrail_information")
|
||||
if not guardrail_information:
|
||||
return
|
||||
|
||||
# _create_guardrail_span reads guardrail_information from
|
||||
# kwargs["standard_logging_object"] and shares its dedupe state via
|
||||
# kwargs["litellm_params"]["metadata"]["_otel_internal"]. Pass the
|
||||
# SAME metadata dict the proxy populated so _handle_failure and
|
||||
# this hook see the same dedupe markers.
|
||||
kwargs: Dict[str, Any] = {
|
||||
"litellm_params": {"metadata": metadata},
|
||||
"standard_logging_object": {
|
||||
"guardrail_information": guardrail_information,
|
||||
"metadata": metadata,
|
||||
},
|
||||
}
|
||||
context = (
|
||||
_trace.set_span_in_context(parent_span) if parent_span is not None else None
|
||||
)
|
||||
self._create_guardrail_span(kwargs=kwargs, context=context)
|
||||
|
||||
async def async_post_call_success_hook(
|
||||
self,
|
||||
data: dict,
|
||||
|
|
@ -736,11 +806,6 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
# Pre-request latency on the SERVER span (success path).
|
||||
self.set_preprocessing_duration_attribute(parent_span, kwargs)
|
||||
|
||||
# http.response.status_code on the SERVER span (success path).
|
||||
# A successful proxy response is HTTP 200; the failure path sets
|
||||
# this from the error code in _record_exception_on_span.
|
||||
self.set_response_status_code_attribute(parent_span, 200)
|
||||
|
||||
# 3. Guardrail span
|
||||
self._create_guardrail_span(kwargs=kwargs, context=ctx)
|
||||
|
||||
|
|
@ -923,7 +988,15 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
and hasattr(proxy_span, "is_recording")
|
||||
and proxy_span.is_recording()
|
||||
):
|
||||
proxy_span.end(end_time=self._to_ns(end_time))
|
||||
self._close_proxy_span_ok(proxy_span, end_time)
|
||||
|
||||
def _close_proxy_span_ok(self, span: Span, end_time) -> None:
|
||||
"""Stamp http.response.status_code=200 + status=OK, then end the span."""
|
||||
from opentelemetry.trace import Status, StatusCode
|
||||
|
||||
self.set_response_status_code_attribute(span, 200)
|
||||
span.set_status(Status(StatusCode.OK))
|
||||
span.end(end_time=self._to_ns(end_time))
|
||||
|
||||
def _handle_success(self, kwargs, response_obj, start_time, end_time):
|
||||
"""Create the litellm_request span then close the proxy span."""
|
||||
|
|
@ -1009,8 +1082,14 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
parent_span is not None
|
||||
and hasattr(parent_span, "name")
|
||||
and parent_span.name == LITELLM_PROXY_REQUEST_SPAN_NAME
|
||||
and hasattr(parent_span, "is_recording")
|
||||
and parent_span.is_recording()
|
||||
):
|
||||
parent_span.end(end_time=self._to_ns(end_time))
|
||||
self._close_proxy_span_ok(parent_span, end_time)
|
||||
|
||||
# Stamp team attributes onto the SERVER (root) span before it is
|
||||
# closed, so the trace root carries them like every child span.
|
||||
self._set_team_attributes_on_proxy_span_from_kwargs(kwargs)
|
||||
|
||||
# close the proxy span explicitly from kwargs metadata
|
||||
# after all child spans (litellm_request, guardrail, raw_request)
|
||||
|
|
@ -1070,8 +1149,70 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
)
|
||||
raw_span.set_status(Status(StatusCode.OK))
|
||||
self.set_raw_request_attributes(raw_span, kwargs, response_obj)
|
||||
self._set_team_attributes_from_kwargs(raw_span, kwargs)
|
||||
raw_span.end(end_time=self._to_ns(end_time))
|
||||
|
||||
def _set_team_attributes_on_span(
|
||||
self,
|
||||
span: Span,
|
||||
team_id: Optional[str],
|
||||
team_alias: Optional[str],
|
||||
) -> None:
|
||||
"""Stamp team_id / team_alias onto a span so every child span of a
|
||||
litellm_request trace carries them, not just the root span.
|
||||
|
||||
Empty strings are treated as absent: a request made with the master
|
||||
key or a team-less virtual key carries ``user_api_key_team_id=""``
|
||||
in ``standard_logging_object.metadata``; propagating that to every
|
||||
span only adds noise that makes traces look mis-instrumented.
|
||||
"""
|
||||
if team_id:
|
||||
self.safe_set_attribute(
|
||||
span=span,
|
||||
key="metadata.user_api_key_team_id",
|
||||
value=team_id,
|
||||
)
|
||||
if team_alias:
|
||||
self.safe_set_attribute(
|
||||
span=span,
|
||||
key="metadata.user_api_key_team_alias",
|
||||
value=team_alias,
|
||||
)
|
||||
|
||||
def _set_team_attributes_from_kwargs(self, span: Span, kwargs: dict) -> None:
|
||||
"""Pull team_id / team_alias from the standard logging metadata in kwargs and stamp them onto span."""
|
||||
std_log = kwargs.get("standard_logging_object")
|
||||
md: dict = {}
|
||||
if isinstance(std_log, dict):
|
||||
md = std_log.get("metadata") or {}
|
||||
elif std_log is not None:
|
||||
md = getattr(std_log, "metadata", None) or {}
|
||||
self._set_team_attributes_on_span(
|
||||
span=span,
|
||||
team_id=md.get("user_api_key_team_id"),
|
||||
team_alias=md.get("user_api_key_team_alias"),
|
||||
)
|
||||
|
||||
def _set_team_attributes_on_proxy_span_from_kwargs(self, kwargs: dict) -> None:
|
||||
"""Stamp team attributes onto the proxy SERVER (root) span so the
|
||||
trace root is filterable by team, not just its children. The root
|
||||
span is created in auth before the team is resolved and is
|
||||
otherwise only closed (never re-attributed) on the success path.
|
||||
|
||||
Guarded to the LiteLLM-created proxy span (by name + recording) so
|
||||
externally provided parent spans are never mutated.
|
||||
"""
|
||||
litellm_params = kwargs.get("litellm_params") or {}
|
||||
metadata = litellm_params.get("metadata") or {}
|
||||
proxy_span = metadata.get("litellm_parent_otel_span")
|
||||
if (
|
||||
proxy_span is not None
|
||||
and getattr(proxy_span, "name", None) == LITELLM_PROXY_REQUEST_SPAN_NAME
|
||||
and hasattr(proxy_span, "is_recording")
|
||||
and proxy_span.is_recording()
|
||||
):
|
||||
self._set_team_attributes_from_kwargs(proxy_span, kwargs)
|
||||
|
||||
def _record_metrics(self, kwargs, response_obj, start_time, end_time):
|
||||
duration_s = (end_time - start_time).total_seconds()
|
||||
params = kwargs.get("litellm_params") or {}
|
||||
|
|
@ -1107,8 +1248,13 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
"mcp_tool_call_metadata",
|
||||
"vector_store_request_metadata",
|
||||
]:
|
||||
if md.get(key) is not None:
|
||||
common_attrs[f"metadata.{key}"] = str(md[key])
|
||||
value = md.get(key)
|
||||
if value is None:
|
||||
continue
|
||||
if isinstance(value, (dict, list)):
|
||||
common_attrs[f"metadata.{key}"] = safe_dumps(value)
|
||||
else:
|
||||
common_attrs[f"metadata.{key}"] = str(value)
|
||||
|
||||
# get hidden params
|
||||
hidden_params = getattr(std_log, "hidden_params", None) or (std_log or {}).get(
|
||||
|
|
@ -1526,12 +1672,45 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
"masked_entity_count", safe_dumps(masked_entity_count)
|
||||
)
|
||||
|
||||
guardrail_response = guardrail_information.get("guardrail_response")
|
||||
if guardrail_response is not None:
|
||||
guardrail_span.set_attribute(
|
||||
"guardrail_response", safe_dumps(guardrail_response)
|
||||
)
|
||||
|
||||
# Surface guardrail_status (success / guardrail_intervened /
|
||||
# guardrail_failed_to_respond / not_run) as a top-level span
|
||||
# attribute so trace backends can filter on it without parsing
|
||||
# guardrail_response.
|
||||
self.safe_set_attribute(
|
||||
span=guardrail_span,
|
||||
key="guardrail_response",
|
||||
value=guardrail_information.get("guardrail_response"),
|
||||
key="guardrail_status",
|
||||
value=guardrail_information.get("guardrail_status"),
|
||||
)
|
||||
|
||||
# Provider's raw top-level action (e.g. Bedrock's
|
||||
# ``GUARDRAIL_INTERVENED`` / ``NONE``). Populated by the provider
|
||||
# hook onto StandardLoggingGuardrailInformation so this integration
|
||||
# stays provider-agnostic — we only read a normalised string.
|
||||
guardrail_action = guardrail_information.get("guardrail_action")
|
||||
if guardrail_action:
|
||||
guardrail_span.set_attribute("guardrail_action", guardrail_action)
|
||||
|
||||
# The provider hook (e.g. Bedrock) extracts violation_categories
|
||||
# from the raw response BEFORE redaction and stamps them onto
|
||||
# StandardLoggingGuardrailInformation. Surfacing them here as a
|
||||
# queryable attribute lets dashboards group by violation category
|
||||
# without parsing the redacted guardrail_response blob.
|
||||
violation_categories = guardrail_information.get("violation_categories")
|
||||
if violation_categories:
|
||||
# OTel sequence attributes must be homogeneous primitives;
|
||||
# serialise to JSON once so set_attribute never coerces.
|
||||
guardrail_span.set_attribute(
|
||||
"guardrail_violation_categories", safe_dumps(violation_categories)
|
||||
)
|
||||
|
||||
self._set_team_attributes_from_kwargs(guardrail_span, kwargs)
|
||||
|
||||
guardrail_span.end(end_time=self._to_ns(end_time_datetime))
|
||||
|
||||
def _handle_failure(self, kwargs, response_obj, start_time, end_time):
|
||||
|
|
@ -2875,6 +3054,11 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
management_endpoint_span.set_status(Status(StatusCode.OK))
|
||||
management_endpoint_span.end(end_time=_end_time_ns)
|
||||
|
||||
# The management wrapper has no other hook that closes the SERVER span.
|
||||
self.set_response_status_code_attribute(parent_otel_span, 200)
|
||||
parent_otel_span.set_status(Status(StatusCode.OK))
|
||||
parent_otel_span.end(end_time=_end_time_ns)
|
||||
|
||||
async def async_management_endpoint_failure_hook(
|
||||
self,
|
||||
logging_payload: ManagementEndpointLoggingPayload,
|
||||
|
|
@ -2925,6 +3109,24 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
management_endpoint_span.set_status(Status(StatusCode.ERROR))
|
||||
management_endpoint_span.end(end_time=_end_time_ns)
|
||||
|
||||
# The management wrapper has no other hook that closes the SERVER span.
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
StandardLoggingPayloadSetup,
|
||||
)
|
||||
|
||||
error_information = StandardLoggingPayloadSetup.get_error_information(
|
||||
original_exception=_exception,
|
||||
)
|
||||
parent_otel_span.set_status(Status(StatusCode.ERROR))
|
||||
self._record_exception_on_span(
|
||||
span=parent_otel_span,
|
||||
kwargs={
|
||||
"exception": _exception,
|
||||
"standard_logging_object": {"error_information": error_information},
|
||||
},
|
||||
)
|
||||
parent_otel_span.end(end_time=_end_time_ns)
|
||||
|
||||
def create_litellm_proxy_request_started_span(
|
||||
self,
|
||||
start_time: datetime,
|
||||
|
|
|
|||
|
|
@ -105,7 +105,7 @@ def _remove_nulls(x: Dict[str, Any]) -> Dict[str, Any]:
|
|||
|
||||
|
||||
def get_traces_and_spans_from_payload(
|
||||
payload: List[Dict[str, Any]]
|
||||
payload: List[Dict[str, Any]],
|
||||
) -> Tuple[List[Dict[str, Any]], List[Dict[str, Any]]]:
|
||||
"""
|
||||
Separate traces and spans from payload.
|
||||
|
|
|
|||
|
|
@ -166,6 +166,53 @@ class PrometheusLogger(CustomLogger):
|
|||
labelnames=self.get_labels_for_metric("litellm_output_tokens_metric"),
|
||||
)
|
||||
|
||||
# Token-type detail metrics. These break out cached, cache-creation,
|
||||
# audio and reasoning tokens that providers report inside
|
||||
# prompt_tokens_details / completion_tokens_details on the usage
|
||||
# object. They are sparse (only incremented when the provider
|
||||
# reports a non-zero value) and are additive to the existing
|
||||
# input/output token totals — no breaking change for existing
|
||||
# dashboards built on the totals.
|
||||
self.litellm_input_cached_tokens_metric = self._counter_factory(
|
||||
"litellm_input_cached_tokens_metric",
|
||||
"Provider-side cached input tokens (e.g. OpenAI prompt_tokens_details.cached_tokens, Anthropic cache_read_input_tokens)",
|
||||
labelnames=self.get_labels_for_metric(
|
||||
"litellm_input_cached_tokens_metric"
|
||||
),
|
||||
)
|
||||
|
||||
self.litellm_input_cache_creation_tokens_metric = self._counter_factory(
|
||||
"litellm_input_cache_creation_tokens_metric",
|
||||
"Provider-side input tokens written to prompt cache (e.g. Anthropic cache_creation_input_tokens)",
|
||||
labelnames=self.get_labels_for_metric(
|
||||
"litellm_input_cache_creation_tokens_metric"
|
||||
),
|
||||
)
|
||||
|
||||
self.litellm_input_audio_tokens_metric = self._counter_factory(
|
||||
"litellm_input_audio_tokens_metric",
|
||||
"Audio input tokens reported in prompt_tokens_details.audio_tokens",
|
||||
labelnames=self.get_labels_for_metric(
|
||||
"litellm_input_audio_tokens_metric"
|
||||
),
|
||||
)
|
||||
|
||||
self.litellm_output_reasoning_tokens_metric = self._counter_factory(
|
||||
"litellm_output_reasoning_tokens_metric",
|
||||
"Reasoning tokens reported in completion_tokens_details.reasoning_tokens",
|
||||
labelnames=self.get_labels_for_metric(
|
||||
"litellm_output_reasoning_tokens_metric"
|
||||
),
|
||||
)
|
||||
|
||||
self.litellm_output_audio_tokens_metric = self._counter_factory(
|
||||
"litellm_output_audio_tokens_metric",
|
||||
"Audio output tokens reported in completion_tokens_details.audio_tokens",
|
||||
labelnames=self.get_labels_for_metric(
|
||||
"litellm_output_audio_tokens_metric"
|
||||
),
|
||||
)
|
||||
|
||||
# Remaining Budget for Team
|
||||
self.litellm_remaining_team_budget_metric = self._gauge_factory(
|
||||
"litellm_remaining_team_budget_metric",
|
||||
|
|
@ -1301,6 +1348,101 @@ class PrometheusLogger(CustomLogger):
|
|||
amount=float(standard_logging_payload["completion_tokens"]),
|
||||
)
|
||||
|
||||
# Token-type detail metrics — sparse, only emitted when the provider
|
||||
# reports a non-zero value in usage.prompt_tokens_details /
|
||||
# usage.completion_tokens_details.
|
||||
self._increment_token_detail_metrics(
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
enum_values=enum_values,
|
||||
label_context=label_context,
|
||||
)
|
||||
|
||||
def _increment_token_detail_metrics(
|
||||
self,
|
||||
standard_logging_payload: StandardLoggingPayload,
|
||||
enum_values: UserAPIKeyLabelValues,
|
||||
label_context: Optional[PrometheusLabelFactoryContext] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Increment per-token-type counters from the Usage object that providers
|
||||
attach to the request. The Usage dict is plumbed onto
|
||||
``standard_logging_payload["metadata"]["usage_object"]`` by
|
||||
``get_standard_logging_object_payload``.
|
||||
|
||||
Each counter is only incremented when the underlying value is > 0, so
|
||||
scrape output stays sparse for providers that don't report these
|
||||
details (most non-OpenAI/Anthropic models).
|
||||
"""
|
||||
metadata = standard_logging_payload.get("metadata") or {}
|
||||
usage_object = (
|
||||
metadata.get("usage_object") if isinstance(metadata, dict) else None
|
||||
)
|
||||
if not isinstance(usage_object, dict):
|
||||
return
|
||||
|
||||
prompt_details = usage_object.get("prompt_tokens_details") or {}
|
||||
completion_details = usage_object.get("completion_tokens_details") or {}
|
||||
|
||||
detail_metrics: List[Tuple[Any, DEFINED_PROMETHEUS_METRICS, Any]] = [
|
||||
(
|
||||
self.litellm_input_cached_tokens_metric,
|
||||
"litellm_input_cached_tokens_metric",
|
||||
(
|
||||
prompt_details.get("cached_tokens")
|
||||
if isinstance(prompt_details, dict)
|
||||
else None
|
||||
),
|
||||
),
|
||||
(
|
||||
self.litellm_input_cache_creation_tokens_metric,
|
||||
"litellm_input_cache_creation_tokens_metric",
|
||||
(
|
||||
prompt_details.get("cache_creation_tokens")
|
||||
if isinstance(prompt_details, dict)
|
||||
else None
|
||||
),
|
||||
),
|
||||
(
|
||||
self.litellm_input_audio_tokens_metric,
|
||||
"litellm_input_audio_tokens_metric",
|
||||
(
|
||||
prompt_details.get("audio_tokens")
|
||||
if isinstance(prompt_details, dict)
|
||||
else None
|
||||
),
|
||||
),
|
||||
(
|
||||
self.litellm_output_reasoning_tokens_metric,
|
||||
"litellm_output_reasoning_tokens_metric",
|
||||
(
|
||||
completion_details.get("reasoning_tokens")
|
||||
if isinstance(completion_details, dict)
|
||||
else None
|
||||
),
|
||||
),
|
||||
(
|
||||
self.litellm_output_audio_tokens_metric,
|
||||
"litellm_output_audio_tokens_metric",
|
||||
(
|
||||
completion_details.get("audio_tokens")
|
||||
if isinstance(completion_details, dict)
|
||||
else None
|
||||
),
|
||||
),
|
||||
]
|
||||
|
||||
for counter, metric_name, value in detail_metrics:
|
||||
if not isinstance(value, (int, float)) or value <= 0:
|
||||
continue
|
||||
PrometheusLogger._inc_labeled_counter(
|
||||
self,
|
||||
counter,
|
||||
metric_name,
|
||||
enum_values,
|
||||
label_context=label_context,
|
||||
amount=float(value),
|
||||
)
|
||||
|
||||
def _increment_cache_metrics(
|
||||
self,
|
||||
standard_logging_payload: StandardLoggingPayload,
|
||||
|
|
@ -3540,6 +3682,10 @@ class PrometheusLogger(CustomLogger):
|
|||
user_object.budget_reset_at = user_info.budget_reset_at
|
||||
if user_object.max_budget is None and user_info.max_budget is not None:
|
||||
user_object.max_budget = user_info.max_budget
|
||||
if user_info.user_email is not None:
|
||||
user_object.user_email = user_info.user_email
|
||||
if user_info.user_alias is not None:
|
||||
user_object.user_alias = user_info.user_alias
|
||||
|
||||
return user_object
|
||||
|
||||
|
|
@ -3556,6 +3702,8 @@ class PrometheusLogger(CustomLogger):
|
|||
"""
|
||||
enum_values = UserAPIKeyLabelValues(
|
||||
user=user.user_id,
|
||||
user_email=user.user_email or "",
|
||||
user_alias=user.user_alias or "",
|
||||
)
|
||||
|
||||
_labels = prometheus_label_factory(
|
||||
|
|
|
|||
605
litellm/integrations/rubrik.py
Normal file
605
litellm/integrations/rubrik.py
Normal file
|
|
@ -0,0 +1,605 @@
|
|||
"""Rubrik LiteLLM Plugin for tool blocking and batch logging."""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import random
|
||||
import time
|
||||
import urllib.parse
|
||||
import uuid
|
||||
from collections import Counter
|
||||
from typing import TYPE_CHECKING, Any, Literal, Optional
|
||||
|
||||
import httpx
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_batch_logger import CustomBatchLogger
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
ModifyResponseException,
|
||||
)
|
||||
from litellm.litellm_core_utils.core_helpers import safe_deep_copy
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import (
|
||||
ChatCompletionMessageToolCall,
|
||||
Function,
|
||||
GenericGuardrailAPIInputs,
|
||||
StandardLoggingPayload,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
Logging as LiteLLMLoggingObj,
|
||||
)
|
||||
|
||||
_ENDPOINT_ANTHROPIC_MESSAGES = "/v1/messages"
|
||||
_WEBHOOK_PATH_TOOL_BLOCKING = "/v1/after_completion/openai/v1"
|
||||
_WEBHOOK_PATH_LOGGING_BATCH = "/v1/litellm/batch"
|
||||
_MAX_QUEUE_SIZE = 10_000
|
||||
_DROP_WARNING_INTERVAL_SECONDS = 60.0
|
||||
|
||||
|
||||
class _MalformedToolBlockingResponseError(Exception):
|
||||
"""Raised when the tool blocking service returns a structurally invalid
|
||||
response (e.g. empty ``choices``).
|
||||
|
||||
Distinct from transient network/HTTP errors so callers can surface a
|
||||
louder, misconfiguration-style log instead of treating it as a routine
|
||||
fail-open.
|
||||
"""
|
||||
|
||||
|
||||
class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
self.flush_lock = asyncio.Lock()
|
||||
kwargs.setdefault("guardrail_name", "rubrik")
|
||||
# `initialize_guardrail` always passes these kwargs explicitly, with
|
||||
# value `None` when the user omits `mode` / `default_on` from the
|
||||
# guardrail config. Coerce None (omitted) to the desired default
|
||||
# while preserving any explicit value the caller did set --
|
||||
# in particular `default_on=False` if the user wants the guardrail
|
||||
# off by default.
|
||||
kwargs["event_hook"] = kwargs.get("event_hook") or GuardrailEventHooks.post_call
|
||||
if kwargs.get("default_on") is None:
|
||||
kwargs["default_on"] = True
|
||||
super().__init__(
|
||||
flush_lock=self.flush_lock,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
verbose_logger.debug("initializing rubrik logger")
|
||||
|
||||
self.sampling_rate = 1.0
|
||||
rbrk_sampling_rate = os.getenv("RUBRIK_SAMPLING_RATE")
|
||||
if rbrk_sampling_rate is not None:
|
||||
try:
|
||||
parsed_rate = float(rbrk_sampling_rate.strip())
|
||||
self.sampling_rate = max(0.0, min(1.0, parsed_rate))
|
||||
if parsed_rate != self.sampling_rate:
|
||||
verbose_logger.warning(
|
||||
f"RUBRIK_SAMPLING_RATE={parsed_rate} clamped to "
|
||||
f"{self.sampling_rate}"
|
||||
)
|
||||
except ValueError:
|
||||
verbose_logger.warning(
|
||||
f"Invalid RUBRIK_SAMPLING_RATE: {rbrk_sampling_rate!r}, using 1.0"
|
||||
)
|
||||
|
||||
self.key = api_key or os.getenv("RUBRIK_API_KEY")
|
||||
if not self.key:
|
||||
verbose_logger.warning(
|
||||
"Rubrik: No API key configured. Requests will be unauthenticated."
|
||||
)
|
||||
_batch_size = os.getenv("RUBRIK_BATCH_SIZE")
|
||||
|
||||
if _batch_size:
|
||||
try:
|
||||
self.batch_size = int(_batch_size)
|
||||
except ValueError:
|
||||
verbose_logger.warning(
|
||||
f"Invalid RUBRIK_BATCH_SIZE: {_batch_size!r}, using default"
|
||||
)
|
||||
|
||||
# Cap the in-memory retry queue so a Rubrik webhook outage cannot let
|
||||
# authenticated traffic accumulate prompt/response payloads until the
|
||||
# proxy runs out of memory. Once the cap is reached, oldest events are
|
||||
# dropped to make room for fresh ones (drop-oldest backpressure).
|
||||
self.max_queue_size = _MAX_QUEUE_SIZE
|
||||
self._dropped_since_warning = 0
|
||||
self._last_drop_warning_time = 0.0
|
||||
|
||||
_webhook_url = api_base or os.getenv("RUBRIK_WEBHOOK_URL")
|
||||
|
||||
if _webhook_url is None:
|
||||
raise ValueError(
|
||||
"Rubrik webhook URL not configured. "
|
||||
"Set RUBRIK_WEBHOOK_URL or pass api_base."
|
||||
)
|
||||
|
||||
_webhook_url = _webhook_url.rstrip("/").removesuffix("/v1")
|
||||
self.tool_blocking_endpoint = f"{_webhook_url}{_WEBHOOK_PATH_TOOL_BLOCKING}"
|
||||
self.logging_endpoint = f"{_webhook_url}{_WEBHOOK_PATH_LOGGING_BATCH}"
|
||||
|
||||
self.async_httpx_client = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.LoggingCallback
|
||||
)
|
||||
|
||||
self.tool_blocking_client = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.LoggingCallback,
|
||||
params={"timeout": httpx.Timeout(5.0, connect=2.0)},
|
||||
)
|
||||
|
||||
self._headers: dict[str, str] = {"Content-Type": "application/json"}
|
||||
if self.key:
|
||||
self._headers["Authorization"] = f"Bearer {self.key}"
|
||||
|
||||
# Periodic flush is started lazily on the first log event so that
|
||||
# low-traffic deployments still get their batches drained even when the
|
||||
# logger is instantiated outside a running event loop (sync init).
|
||||
self._flush_task: Optional[asyncio.Task[Any]] = (
|
||||
self._start_periodic_flush_task()
|
||||
)
|
||||
|
||||
def _start_periodic_flush_task(self) -> Optional[asyncio.Task[Any]]:
|
||||
"""Start the periodic flush task only when an event loop is already running."""
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
verbose_logger.debug(
|
||||
"Rubrik logger init: no running event loop, "
|
||||
"periodic flush will start on first log event."
|
||||
)
|
||||
return None
|
||||
return loop.create_task(self.periodic_flush())
|
||||
|
||||
def _ensure_periodic_flush_task(self) -> None:
|
||||
# Synchronous helper: in asyncio's cooperative model there is no await
|
||||
# between the check and assignment, so two callers cannot race here.
|
||||
if self._flush_task is None or self._flush_task.done():
|
||||
self._flush_task = self._start_periodic_flush_task()
|
||||
|
||||
async def aclose(self):
|
||||
"""Close the dedicated HTTP clients used by this logger."""
|
||||
# Cancel the periodic flush task before closing the HTTP clients so
|
||||
# the loop doesn't wake up and try to POST via a closed client.
|
||||
if self._flush_task is not None and not self._flush_task.done():
|
||||
self._flush_task.cancel()
|
||||
try:
|
||||
await self._flush_task
|
||||
except (asyncio.CancelledError, Exception):
|
||||
pass
|
||||
self._flush_task = None
|
||||
await self.tool_blocking_client.close()
|
||||
await self.async_httpx_client.close()
|
||||
|
||||
# -- Guardrail hook --------------------------------------------------------
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
"""Validate tool calls against the blocking service (fail-open)."""
|
||||
if input_type != "response":
|
||||
return inputs
|
||||
|
||||
tool_calls = inputs.get("tool_calls")
|
||||
if not tool_calls:
|
||||
return inputs
|
||||
|
||||
try:
|
||||
return await self._check_tool_calls(
|
||||
inputs, tool_calls, request_data, logging_obj
|
||||
)
|
||||
except ModifyResponseException:
|
||||
raise
|
||||
except _MalformedToolBlockingResponseError as e:
|
||||
# Distinct from transient errors: the service responded but the
|
||||
# payload was structurally invalid, which usually indicates a
|
||||
# misconfigured webhook or a breaking change in its response
|
||||
# format. Log loudly so operators notice their tool-blocking
|
||||
# policy is not actually being enforced.
|
||||
verbose_logger.critical(
|
||||
"Tool blocking service returned a malformed response: %s. "
|
||||
"Tool calls are NOT being checked -- verify the webhook "
|
||||
"configuration. Returning original response unchanged.",
|
||||
e,
|
||||
exc_info=True,
|
||||
)
|
||||
return inputs
|
||||
except Exception as e:
|
||||
verbose_logger.error(
|
||||
f"Tool blocking hook failed: {e}. "
|
||||
"Returning original response unchanged.",
|
||||
exc_info=True,
|
||||
)
|
||||
return inputs
|
||||
|
||||
async def _check_tool_calls(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
tool_calls: Any,
|
||||
request_data: dict,
|
||||
logging_obj: Optional["LiteLLMLoggingObj"],
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
"""Send tool calls to blocking service, raise if any are blocked."""
|
||||
message_tool_calls = self._normalize_tool_calls(tool_calls)
|
||||
|
||||
call_details = (
|
||||
getattr(logging_obj, "model_call_details", {}) if logging_obj else {}
|
||||
)
|
||||
response = request_data.get("response")
|
||||
request_id = getattr(response, "id", None) if response else None
|
||||
if logging_obj and not call_details:
|
||||
verbose_logger.warning(
|
||||
"Rubrik: logging_obj present but model_call_details is empty "
|
||||
"-- request context will be missing"
|
||||
)
|
||||
|
||||
response_data = self._build_tool_call_payload(message_tool_calls, request_id)
|
||||
req_data = self._extract_request_data(call_details)
|
||||
|
||||
service_response = await self._post_to_tool_blocking_service(
|
||||
response_data, req_data
|
||||
)
|
||||
blocked_explanation = self._extract_blocked_tools(
|
||||
service_response, message_tool_calls
|
||||
)
|
||||
|
||||
if blocked_explanation is not None:
|
||||
model = self._resolve_model(request_data, call_details)
|
||||
raise ModifyResponseException(
|
||||
message=blocked_explanation,
|
||||
model=model,
|
||||
request_data=request_data,
|
||||
guardrail_name=self.guardrail_name,
|
||||
)
|
||||
|
||||
return inputs
|
||||
|
||||
@staticmethod
|
||||
def _normalize_tool_calls(tool_calls: Any) -> list[ChatCompletionMessageToolCall]:
|
||||
"""Convert tool_calls from inputs to ChatCompletionMessageToolCall objects."""
|
||||
result = []
|
||||
for tc in tool_calls:
|
||||
if isinstance(tc, ChatCompletionMessageToolCall):
|
||||
result.append(tc)
|
||||
elif isinstance(tc, dict):
|
||||
func = tc.get("function", {})
|
||||
result.append(
|
||||
ChatCompletionMessageToolCall(
|
||||
id=tc.get("id", ""),
|
||||
type=tc.get("type", "function"),
|
||||
function=Function(
|
||||
name=func.get("name", ""),
|
||||
arguments=func.get("arguments", ""),
|
||||
),
|
||||
)
|
||||
)
|
||||
elif hasattr(tc, "id") and hasattr(tc, "function"):
|
||||
result.append(
|
||||
ChatCompletionMessageToolCall(
|
||||
id=tc.id or "",
|
||||
type=getattr(tc, "type", None) or "function",
|
||||
function=tc.function,
|
||||
)
|
||||
)
|
||||
else:
|
||||
raise TypeError(
|
||||
f"Cannot normalize tool_call of type {type(tc).__name__}"
|
||||
)
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _build_tool_call_payload(
|
||||
tool_calls: list[ChatCompletionMessageToolCall],
|
||||
request_id: str | None,
|
||||
) -> dict[str, Any]:
|
||||
"""Build a full OpenAI ChatCompletion-format dict for the blocking service."""
|
||||
return {
|
||||
"id": request_id or f"chatcmpl-{uuid.uuid4()}",
|
||||
"object": "chat.completion",
|
||||
"created": int(time.time()),
|
||||
"model": "",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
tc.model_dump(exclude_none=True) for tc in tool_calls
|
||||
],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _extract_request_data(call_details: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Extract original request data from model_call_details."""
|
||||
if not call_details:
|
||||
return {}
|
||||
litellm_params = call_details.get("litellm_params", {}) or {}
|
||||
return {
|
||||
"messages": call_details.get("messages"),
|
||||
"model": call_details.get("model"),
|
||||
"proxy_server_request": RubrikLogger._sanitize_proxy_server_request(
|
||||
litellm_params.get("proxy_server_request")
|
||||
),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _sanitize_proxy_server_request(proxy_server_request: Any) -> Any:
|
||||
"""Allowlist only routing fields (``url``, ``method``) when forwarding
|
||||
``proxy_server_request`` to the external Rubrik webhook, dropping
|
||||
inbound ``headers`` (Authorization, Cookie, x-api-key, ...) and the raw
|
||||
request ``body`` so proxy credentials are not exfiltrated."""
|
||||
if not isinstance(proxy_server_request, dict):
|
||||
return proxy_server_request
|
||||
return {
|
||||
key: proxy_server_request[key]
|
||||
for key in ("url", "method")
|
||||
if key in proxy_server_request
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _resolve_model(
|
||||
request_data: dict[str, Any], call_details: dict[str, Any]
|
||||
) -> str:
|
||||
"""Get the model name for the ModifyResponseException."""
|
||||
response = request_data.get("response")
|
||||
if response and hasattr(response, "model"):
|
||||
return response.model or "unknown"
|
||||
return call_details.get("model", "unknown")
|
||||
|
||||
# -- Logging hooks ---------------------------------------------------------
|
||||
|
||||
async def _prepare_log_payload(
|
||||
self, kwargs: dict, event_type: str
|
||||
) -> StandardLoggingPayload | None:
|
||||
"""Shared logic for success and failure logging."""
|
||||
if random.random() > self.sampling_rate:
|
||||
verbose_logger.debug(
|
||||
f"Skipping Rubrik {event_type} logging "
|
||||
f"(sampling_rate={self.sampling_rate})"
|
||||
)
|
||||
return None
|
||||
|
||||
# Deep-copy so mutations don't affect other callbacks sharing this object
|
||||
standard_logging_payload: StandardLoggingPayload = safe_deep_copy(
|
||||
kwargs["standard_logging_object"]
|
||||
)
|
||||
|
||||
# For Anthropic /v1/messages requests, LiteLLM creates a separate
|
||||
# ModelResponse (with a generated chatcmpl-* id) for logging, which
|
||||
# differs from the original Anthropic msg-* id on the response dict.
|
||||
# Normalize to litellm_call_id so that the logging and tool-blocking
|
||||
# endpoints see the same request identifier.
|
||||
litellm_params = kwargs.get("litellm_params", {}) or {}
|
||||
proxy_request = litellm_params.get("proxy_server_request", {}) or {}
|
||||
url_path = urllib.parse.urlparse(proxy_request.get("url", "")).path
|
||||
if url_path.endswith(_ENDPOINT_ANTHROPIC_MESSAGES):
|
||||
_litellm_call_id = kwargs.get("litellm_call_id")
|
||||
if _litellm_call_id:
|
||||
standard_logging_payload["id"] = _litellm_call_id # type: ignore[literal-required]
|
||||
|
||||
if "system" in kwargs:
|
||||
system_prompt_msg_list = kwargs["system"]
|
||||
try:
|
||||
if system_prompt_msg_list:
|
||||
system_scaffold = {
|
||||
"role": "system",
|
||||
"content": system_prompt_msg_list,
|
||||
}
|
||||
if isinstance(standard_logging_payload["messages"], list):
|
||||
standard_logging_payload["messages"].insert(0, system_scaffold)
|
||||
elif isinstance(standard_logging_payload["messages"], (dict, str)):
|
||||
standard_logging_payload["messages"] = [
|
||||
system_scaffold,
|
||||
standard_logging_payload["messages"],
|
||||
]
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
f"Rubrik: failed to prepend system prompt: {e}",
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
return standard_logging_payload
|
||||
|
||||
async def _enqueue_log_event(self, kwargs: dict, event_type: str):
|
||||
try:
|
||||
self._ensure_periodic_flush_task()
|
||||
payload = await self._prepare_log_payload(kwargs, event_type)
|
||||
if payload is None:
|
||||
return
|
||||
|
||||
self.log_queue.append(payload)
|
||||
self._enforce_max_queue_size()
|
||||
|
||||
if len(self.log_queue) >= self.batch_size:
|
||||
await self.flush_queue()
|
||||
except Exception as e:
|
||||
verbose_logger.error(
|
||||
f"Rubrik {event_type} logging hook failed: {e}. "
|
||||
"Skipping logging for this event.",
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
def _enforce_max_queue_size(self) -> None:
|
||||
overflow = len(self.log_queue) - self.max_queue_size
|
||||
if overflow <= 0:
|
||||
return
|
||||
del self.log_queue[:overflow]
|
||||
self._dropped_since_warning += overflow
|
||||
now = time.time()
|
||||
if now - self._last_drop_warning_time >= _DROP_WARNING_INTERVAL_SECONDS:
|
||||
verbose_logger.warning(
|
||||
"Rubrik: log queue exceeded max_queue_size=%s; dropped %s "
|
||||
"oldest events since the last warning. The Rubrik webhook may "
|
||||
"be unhealthy or undersized for current traffic.",
|
||||
self.max_queue_size,
|
||||
self._dropped_since_warning,
|
||||
)
|
||||
self._dropped_since_warning = 0
|
||||
self._last_drop_warning_time = now
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
await self._enqueue_log_event(kwargs, "success")
|
||||
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
await self._enqueue_log_event(kwargs, "failure")
|
||||
|
||||
# -- Batch logging ---------------------------------------------------------
|
||||
|
||||
async def _log_batch_to_rubrik(self, data):
|
||||
# NOTE: this method intentionally re-raises on failure so the parent
|
||||
# CustomBatchLogger.flush_queue keeps the unsent events in the queue
|
||||
# for the next flush attempt instead of silently dropping them.
|
||||
try:
|
||||
response = await self.async_httpx_client.post(
|
||||
url=self.logging_endpoint,
|
||||
json=data,
|
||||
headers=self._headers,
|
||||
)
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
verbose_logger.exception(
|
||||
f"Rubrik HTTP Error: {e.response.status_code} - {e.response.text}"
|
||||
)
|
||||
raise
|
||||
except Exception:
|
||||
verbose_logger.exception("Rubrik Layer Error")
|
||||
raise
|
||||
|
||||
async def async_send_batch(self):
|
||||
"""Handles sending batches of responses to Rubrik.
|
||||
|
||||
Note: the canonical flush path is :meth:`flush_queue`, which takes a
|
||||
single snapshot used for both sending and queue draining. This method
|
||||
is kept for direct callers / tests; it intentionally does NOT remove
|
||||
events from the queue.
|
||||
"""
|
||||
if not self.log_queue:
|
||||
return
|
||||
|
||||
log_queue_snapshot = list(self.log_queue)
|
||||
verbose_logger.debug(
|
||||
"Rubrik: Flushing batch of %s events", len(log_queue_snapshot)
|
||||
)
|
||||
await self._log_batch_to_rubrik(
|
||||
data=log_queue_snapshot,
|
||||
)
|
||||
|
||||
async def flush_queue(self):
|
||||
"""Snapshot, send, and drain in one consistent step.
|
||||
|
||||
Overrides the base implementation so the same snapshot drives both
|
||||
the HTTP send and the queue truncation. This avoids the subtle
|
||||
coupling where the base class captures `len(self.log_queue)`
|
||||
separately from the snapshot taken inside `async_send_batch`,
|
||||
which could otherwise drift in a future refactor and cause
|
||||
duplicate deliveries to Rubrik.
|
||||
"""
|
||||
if self.flush_lock is None:
|
||||
return
|
||||
|
||||
async with self.flush_lock:
|
||||
if not self.log_queue:
|
||||
return
|
||||
snapshot = list(self.log_queue)
|
||||
verbose_logger.debug("Rubrik: Flushing batch of %s events", len(snapshot))
|
||||
try:
|
||||
await self._log_batch_to_rubrik(data=snapshot)
|
||||
except Exception:
|
||||
# Already logged with traceback inside _log_batch_to_rubrik.
|
||||
# Preserve the in-flight events for retry on the next flush.
|
||||
return
|
||||
del self.log_queue[: len(snapshot)]
|
||||
self.last_flush_time = time.time()
|
||||
|
||||
# -- Tool blocking service -------------------------------------------------
|
||||
|
||||
async def _post_to_tool_blocking_service(
|
||||
self,
|
||||
response_data: dict[str, Any],
|
||||
request_data: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""Post a payload to the tool blocking service and return the response.
|
||||
|
||||
Args:
|
||||
response_data: The OpenAI-formatted response payload to send.
|
||||
request_data: Original LLM request data to include alongside
|
||||
the response for additional context. Empty dict if unavailable.
|
||||
|
||||
Raises:
|
||||
Exception: If the service is unavailable or returns an error.
|
||||
"""
|
||||
envelope = {
|
||||
"request": request_data,
|
||||
"response": response_data,
|
||||
}
|
||||
verbose_logger.debug(
|
||||
f"Sending request to tool blocking service: "
|
||||
f"{self.tool_blocking_endpoint}"
|
||||
)
|
||||
http_response = await self.tool_blocking_client.post(
|
||||
self.tool_blocking_endpoint,
|
||||
json=envelope,
|
||||
headers=self._headers,
|
||||
)
|
||||
http_response.raise_for_status()
|
||||
result: dict[str, Any] = http_response.json()
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _extract_blocked_tools(
|
||||
service_response: dict[str, Any],
|
||||
all_tool_calls: list[ChatCompletionMessageToolCall],
|
||||
) -> Optional[str]:
|
||||
"""Return the blocking explanation if any tool calls were blocked.
|
||||
|
||||
Compares the service response (which contains only allowed tools) against
|
||||
the full set of tool calls. Returns ``None`` if all tools are allowed, or
|
||||
the explanation string (prefixed with newlines) otherwise.
|
||||
|
||||
Expects service_response in OpenAI chat completion format:
|
||||
{"choices": [{"message": {"tool_calls": [...], "content": "..."}}]}
|
||||
"""
|
||||
choices = service_response.get("choices", [])
|
||||
if not choices:
|
||||
raise _MalformedToolBlockingResponseError(
|
||||
"Tool blocking service returned empty response"
|
||||
)
|
||||
|
||||
message = choices[0].get("message", {})
|
||||
returned_tool_calls = message.get("tool_calls") or []
|
||||
blocking_explanation = message.get("content", "")
|
||||
|
||||
allowed_id_counts: Counter = Counter(
|
||||
tc["id"]
|
||||
for tc in returned_tool_calls
|
||||
if isinstance(tc, dict) and tc.get("id")
|
||||
)
|
||||
required_id_counts: Counter = Counter(tc.id for tc in all_tool_calls if tc.id)
|
||||
|
||||
all_allowed = len(returned_tool_calls) >= len(all_tool_calls) and all(
|
||||
allowed_id_counts.get(tc_id, 0) >= count
|
||||
for tc_id, count in required_id_counts.items()
|
||||
)
|
||||
|
||||
if all_allowed:
|
||||
return None
|
||||
|
||||
explanation = blocking_explanation or "Tool call blocked by policy."
|
||||
return f"\n\n{explanation}"
|
||||
|
|
@ -1,8 +1,8 @@
|
|||
"""
|
||||
s3 Bucket Logging Integration
|
||||
|
||||
async_log_success_event: Processes the event, stores it in memory for DEFAULT_S3_FLUSH_INTERVAL_SECONDS seconds or until DEFAULT_S3_BATCH_SIZE and then flushes to s3
|
||||
async_log_failure_event: Processes the event, stores it in memory for DEFAULT_S3_FLUSH_INTERVAL_SECONDS seconds or until DEFAULT_S3_BATCH_SIZE and then flushes to s3
|
||||
async_log_success_event: Processes the event, stores it in memory for DEFAULT_S3_FLUSH_INTERVAL_SECONDS seconds or until DEFAULT_S3_BATCH_SIZE and then flushes to s3
|
||||
async_log_failure_event: Processes the event, stores it in memory for DEFAULT_S3_FLUSH_INTERVAL_SECONDS seconds or until DEFAULT_S3_BATCH_SIZE and then flushes to s3
|
||||
NOTE 1: S3 does not provide a BATCH PUT API endpoint, so we create tasks to upload each element individually
|
||||
"""
|
||||
|
||||
|
|
|
|||
|
|
@ -5,31 +5,40 @@ This module provides SDK methods for Google's Interactions API.
|
|||
|
||||
Usage:
|
||||
import litellm
|
||||
|
||||
|
||||
# Create an interaction with a model
|
||||
response = litellm.interactions.create(
|
||||
model="gemini-2.5-flash",
|
||||
input="Hello, how are you?"
|
||||
)
|
||||
|
||||
|
||||
# Create an interaction with an agent
|
||||
response = litellm.interactions.create(
|
||||
agent="deep-research-pro-preview-12-2025",
|
||||
input="Research the current state of cancer research"
|
||||
)
|
||||
|
||||
|
||||
# Async version
|
||||
response = await litellm.interactions.acreate(...)
|
||||
|
||||
|
||||
# Get an interaction
|
||||
response = litellm.interactions.get(interaction_id="...")
|
||||
|
||||
|
||||
# Delete an interaction
|
||||
result = litellm.interactions.delete(interaction_id="...")
|
||||
|
||||
|
||||
# Cancel an interaction
|
||||
result = litellm.interactions.cancel(interaction_id="...")
|
||||
|
||||
# Create a managed agent on the provider side
|
||||
result = litellm.interactions.agents.create(
|
||||
name="waverunner",
|
||||
custom_llm_provider="gemini",
|
||||
api_key="...",
|
||||
base_agent="gemini-2.5-flash",
|
||||
instructions="You are a helpful assistant.",
|
||||
)
|
||||
|
||||
Methods:
|
||||
- create(): Sync create interaction
|
||||
- acreate(): Async create interaction
|
||||
|
|
@ -39,8 +48,12 @@ Methods:
|
|||
- adelete(): Async delete interaction
|
||||
- cancel(): Sync cancel interaction
|
||||
- acancel(): Async cancel interaction
|
||||
|
||||
Sub-modules:
|
||||
- agents: Provider-side agent creation (litellm.interactions.agents.create)
|
||||
"""
|
||||
|
||||
from litellm.interactions import agents
|
||||
from litellm.interactions.main import (
|
||||
acancel,
|
||||
acreate,
|
||||
|
|
@ -65,4 +78,6 @@ __all__ = [
|
|||
# Cancel
|
||||
"cancel",
|
||||
"acancel",
|
||||
# Sub-modules
|
||||
"agents",
|
||||
]
|
||||
|
|
|
|||
39
litellm/interactions/agents/__init__.py
Normal file
39
litellm/interactions/agents/__init__.py
Normal file
|
|
@ -0,0 +1,39 @@
|
|||
"""
|
||||
litellm.interactions.agents
|
||||
|
||||
Full CRUD SDK for provider-side managed agents (e.g. Gemini v1beta/agents).
|
||||
|
||||
litellm.interactions.agents.create(name=..., ...)
|
||||
litellm.interactions.agents.list(api_key=...)
|
||||
litellm.interactions.agents.get(name=..., ...)
|
||||
litellm.interactions.agents.delete(name=..., ...)
|
||||
litellm.interactions.agents.list_versions(name=..., ...)
|
||||
|
||||
Async counterparts: acreate, alist, aget, adelete, alist_versions
|
||||
"""
|
||||
|
||||
from litellm.interactions.agents.main import (
|
||||
acreate,
|
||||
adelete,
|
||||
aget,
|
||||
alist,
|
||||
alist_versions,
|
||||
create,
|
||||
delete,
|
||||
get,
|
||||
list,
|
||||
list_versions,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"create",
|
||||
"acreate",
|
||||
"list",
|
||||
"alist",
|
||||
"get",
|
||||
"aget",
|
||||
"delete",
|
||||
"adelete",
|
||||
"list_versions",
|
||||
"alist_versions",
|
||||
]
|
||||
478
litellm/interactions/agents/http_handler.py
Normal file
478
litellm/interactions/agents/http_handler.py
Normal file
|
|
@ -0,0 +1,478 @@
|
|||
"""
|
||||
HTTP handler for the Agents API.
|
||||
|
||||
Extends InteractionsHTTPHandler so that the shared HTTP infrastructure
|
||||
(_handle_error, _sync_client, _async_client) is reused rather than
|
||||
duplicated. BaseAgentsAPIConfig stays as pure transform code.
|
||||
"""
|
||||
|
||||
from typing import Any, Coroutine, Dict, Optional, Union
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.constants import request_timeout
|
||||
from litellm.interactions.http_handler import InteractionsHTTPHandler
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.agents.transformation import BaseAgentsAPIConfig
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.types.agents import (
|
||||
AgentCreateResponse,
|
||||
AgentDeleteResult,
|
||||
AgentListResponse,
|
||||
AgentVersionsResponse,
|
||||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
|
||||
class AgentsHTTPHandler(InteractionsHTTPHandler):
|
||||
"""HTTP handler for Agents API CRUD requests."""
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# CREATE #
|
||||
# ------------------------------------------------------------------ #
|
||||
|
||||
def create_agent(
|
||||
self,
|
||||
agents_api_config: BaseAgentsAPIConfig,
|
||||
name: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
client: Optional[HTTPHandler] = None,
|
||||
_is_async: bool = False,
|
||||
) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]:
|
||||
if _is_async:
|
||||
return self.async_create_agent(
|
||||
agents_api_config=agents_api_config,
|
||||
name=name,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
extra_body=extra_body,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
sync_httpx_client = self._sync_client(litellm_params, client)
|
||||
headers = agents_api_config.validate_environment(
|
||||
headers=extra_headers or {}, litellm_params=dict(litellm_params)
|
||||
)
|
||||
url = agents_api_config.get_complete_url(
|
||||
api_base=litellm_params.get("api_base"),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
data = agents_api_config.transform_create_request(
|
||||
name=name, litellm_params=dict(litellm_params)
|
||||
)
|
||||
if extra_body:
|
||||
data.update(extra_body)
|
||||
|
||||
logging_obj.pre_call(
|
||||
input=name,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": url,
|
||||
"headers": headers,
|
||||
},
|
||||
)
|
||||
try:
|
||||
response = sync_httpx_client.post(
|
||||
url=url, headers=headers, json=data, timeout=timeout or request_timeout
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=agents_api_config)
|
||||
|
||||
logging_obj.post_call(
|
||||
original_response=response.text,
|
||||
additional_args={"complete_input_dict": data},
|
||||
)
|
||||
return agents_api_config.transform_create_response(
|
||||
raw_response=response, name=name
|
||||
)
|
||||
|
||||
async def async_create_agent(
|
||||
self,
|
||||
agents_api_config: BaseAgentsAPIConfig,
|
||||
name: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
client: Optional[AsyncHTTPHandler] = None,
|
||||
) -> AgentCreateResponse:
|
||||
async_httpx_client = self._async_client(litellm_params, client)
|
||||
headers = agents_api_config.validate_environment(
|
||||
headers=extra_headers or {}, litellm_params=dict(litellm_params)
|
||||
)
|
||||
url = agents_api_config.get_complete_url(
|
||||
api_base=litellm_params.get("api_base"),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
data = agents_api_config.transform_create_request(
|
||||
name=name, litellm_params=dict(litellm_params)
|
||||
)
|
||||
if extra_body:
|
||||
data.update(extra_body)
|
||||
|
||||
logging_obj.pre_call(
|
||||
input=name,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": url,
|
||||
"headers": headers,
|
||||
},
|
||||
)
|
||||
try:
|
||||
response = await async_httpx_client.post(
|
||||
url=url, headers=headers, json=data, timeout=timeout or request_timeout
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=agents_api_config)
|
||||
|
||||
logging_obj.post_call(
|
||||
original_response=response.text,
|
||||
additional_args={"complete_input_dict": data},
|
||||
)
|
||||
return agents_api_config.transform_create_response(
|
||||
raw_response=response, name=name
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# LIST #
|
||||
# ------------------------------------------------------------------ #
|
||||
|
||||
def list_agents(
|
||||
self,
|
||||
agents_api_config: BaseAgentsAPIConfig,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
client: Optional[HTTPHandler] = None,
|
||||
_is_async: bool = False,
|
||||
) -> Union[AgentListResponse, Coroutine[Any, Any, AgentListResponse]]:
|
||||
if _is_async:
|
||||
return self.async_list_agents(
|
||||
agents_api_config=agents_api_config,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
sync_httpx_client = self._sync_client(litellm_params, client)
|
||||
headers = agents_api_config.validate_environment(
|
||||
headers=extra_headers or {}, litellm_params=dict(litellm_params)
|
||||
)
|
||||
url, params = agents_api_config.transform_list_request(
|
||||
api_base=litellm_params.get("api_base"),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
logging_obj.pre_call(
|
||||
input="list_agents",
|
||||
api_key="",
|
||||
additional_args={"api_base": url, "headers": headers},
|
||||
)
|
||||
try:
|
||||
response = sync_httpx_client.get(url=url, headers=headers, params=params)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=agents_api_config)
|
||||
|
||||
logging_obj.post_call(original_response=response.text, additional_args={})
|
||||
return agents_api_config.transform_list_response(raw_response=response)
|
||||
|
||||
async def async_list_agents(
|
||||
self,
|
||||
agents_api_config: BaseAgentsAPIConfig,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
client: Optional[AsyncHTTPHandler] = None,
|
||||
) -> AgentListResponse:
|
||||
async_httpx_client = self._async_client(litellm_params, client)
|
||||
headers = agents_api_config.validate_environment(
|
||||
headers=extra_headers or {}, litellm_params=dict(litellm_params)
|
||||
)
|
||||
url, params = agents_api_config.transform_list_request(
|
||||
api_base=litellm_params.get("api_base"),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
logging_obj.pre_call(
|
||||
input="list_agents",
|
||||
api_key="",
|
||||
additional_args={"api_base": url, "headers": headers},
|
||||
)
|
||||
try:
|
||||
response = await async_httpx_client.get(
|
||||
url=url, headers=headers, params=params
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=agents_api_config)
|
||||
|
||||
logging_obj.post_call(original_response=response.text, additional_args={})
|
||||
return agents_api_config.transform_list_response(raw_response=response)
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# GET #
|
||||
# ------------------------------------------------------------------ #
|
||||
|
||||
def get_agent(
|
||||
self,
|
||||
agents_api_config: BaseAgentsAPIConfig,
|
||||
name: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
client: Optional[HTTPHandler] = None,
|
||||
_is_async: bool = False,
|
||||
) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]:
|
||||
if _is_async:
|
||||
return self.async_get_agent(
|
||||
agents_api_config=agents_api_config,
|
||||
name=name,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
sync_httpx_client = self._sync_client(litellm_params, client)
|
||||
headers = agents_api_config.validate_environment(
|
||||
headers=extra_headers or {}, litellm_params=dict(litellm_params)
|
||||
)
|
||||
url, params = agents_api_config.transform_get_request(
|
||||
name=name,
|
||||
api_base=litellm_params.get("api_base"),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
logging_obj.pre_call(
|
||||
input=name,
|
||||
api_key="",
|
||||
additional_args={"api_base": url, "headers": headers},
|
||||
)
|
||||
try:
|
||||
response = sync_httpx_client.get(url=url, headers=headers, params=params)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=agents_api_config)
|
||||
|
||||
logging_obj.post_call(original_response=response.text, additional_args={})
|
||||
return agents_api_config.transform_get_response(
|
||||
raw_response=response, name=name
|
||||
)
|
||||
|
||||
async def async_get_agent(
|
||||
self,
|
||||
agents_api_config: BaseAgentsAPIConfig,
|
||||
name: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
client: Optional[AsyncHTTPHandler] = None,
|
||||
) -> AgentCreateResponse:
|
||||
async_httpx_client = self._async_client(litellm_params, client)
|
||||
headers = agents_api_config.validate_environment(
|
||||
headers=extra_headers or {}, litellm_params=dict(litellm_params)
|
||||
)
|
||||
url, params = agents_api_config.transform_get_request(
|
||||
name=name,
|
||||
api_base=litellm_params.get("api_base"),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
logging_obj.pre_call(
|
||||
input=name,
|
||||
api_key="",
|
||||
additional_args={"api_base": url, "headers": headers},
|
||||
)
|
||||
try:
|
||||
response = await async_httpx_client.get(
|
||||
url=url, headers=headers, params=params
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=agents_api_config)
|
||||
|
||||
logging_obj.post_call(original_response=response.text, additional_args={})
|
||||
return agents_api_config.transform_get_response(
|
||||
raw_response=response, name=name
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# DELETE #
|
||||
# ------------------------------------------------------------------ #
|
||||
|
||||
def delete_agent(
|
||||
self,
|
||||
agents_api_config: BaseAgentsAPIConfig,
|
||||
name: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
client: Optional[HTTPHandler] = None,
|
||||
_is_async: bool = False,
|
||||
) -> Union[AgentDeleteResult, Coroutine[Any, Any, AgentDeleteResult]]:
|
||||
if _is_async:
|
||||
return self.async_delete_agent(
|
||||
agents_api_config=agents_api_config,
|
||||
name=name,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
sync_httpx_client = self._sync_client(litellm_params, client)
|
||||
headers = agents_api_config.validate_environment(
|
||||
headers=extra_headers or {}, litellm_params=dict(litellm_params)
|
||||
)
|
||||
url = agents_api_config.transform_delete_request(
|
||||
name=name,
|
||||
api_base=litellm_params.get("api_base"),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
logging_obj.pre_call(
|
||||
input=name,
|
||||
api_key="",
|
||||
additional_args={"api_base": url, "headers": headers},
|
||||
)
|
||||
try:
|
||||
response = sync_httpx_client.delete(
|
||||
url=url, headers=headers, timeout=timeout or request_timeout
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=agents_api_config)
|
||||
|
||||
logging_obj.post_call(original_response=response.text, additional_args={})
|
||||
return agents_api_config.transform_delete_response(
|
||||
raw_response=response, name=name
|
||||
)
|
||||
|
||||
async def async_delete_agent(
|
||||
self,
|
||||
agents_api_config: BaseAgentsAPIConfig,
|
||||
name: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
client: Optional[AsyncHTTPHandler] = None,
|
||||
) -> AgentDeleteResult:
|
||||
async_httpx_client = self._async_client(litellm_params, client)
|
||||
headers = agents_api_config.validate_environment(
|
||||
headers=extra_headers or {}, litellm_params=dict(litellm_params)
|
||||
)
|
||||
url = agents_api_config.transform_delete_request(
|
||||
name=name,
|
||||
api_base=litellm_params.get("api_base"),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
logging_obj.pre_call(
|
||||
input=name,
|
||||
api_key="",
|
||||
additional_args={"api_base": url, "headers": headers},
|
||||
)
|
||||
try:
|
||||
response = await async_httpx_client.delete(
|
||||
url=url, headers=headers, timeout=timeout or request_timeout
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=agents_api_config)
|
||||
|
||||
logging_obj.post_call(original_response=response.text, additional_args={})
|
||||
return agents_api_config.transform_delete_response(
|
||||
raw_response=response, name=name
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# LIST VERSIONS #
|
||||
# ------------------------------------------------------------------ #
|
||||
|
||||
def list_agent_versions(
|
||||
self,
|
||||
agents_api_config: BaseAgentsAPIConfig,
|
||||
name: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
client: Optional[HTTPHandler] = None,
|
||||
_is_async: bool = False,
|
||||
) -> Union[AgentVersionsResponse, Coroutine[Any, Any, AgentVersionsResponse]]:
|
||||
if _is_async:
|
||||
return self.async_list_agent_versions(
|
||||
agents_api_config=agents_api_config,
|
||||
name=name,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
sync_httpx_client = self._sync_client(litellm_params, client)
|
||||
headers = agents_api_config.validate_environment(
|
||||
headers=extra_headers or {}, litellm_params=dict(litellm_params)
|
||||
)
|
||||
url, params = agents_api_config.transform_list_versions_request(
|
||||
name=name,
|
||||
api_base=litellm_params.get("api_base"),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
logging_obj.pre_call(
|
||||
input=name,
|
||||
api_key="",
|
||||
additional_args={"api_base": url, "headers": headers},
|
||||
)
|
||||
try:
|
||||
response = sync_httpx_client.get(url=url, headers=headers, params=params)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=agents_api_config)
|
||||
|
||||
logging_obj.post_call(original_response=response.text, additional_args={})
|
||||
return agents_api_config.transform_list_versions_response(
|
||||
raw_response=response, name=name
|
||||
)
|
||||
|
||||
async def async_list_agent_versions(
|
||||
self,
|
||||
agents_api_config: BaseAgentsAPIConfig,
|
||||
name: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
client: Optional[AsyncHTTPHandler] = None,
|
||||
) -> AgentVersionsResponse:
|
||||
async_httpx_client = self._async_client(litellm_params, client)
|
||||
headers = agents_api_config.validate_environment(
|
||||
headers=extra_headers or {}, litellm_params=dict(litellm_params)
|
||||
)
|
||||
url, params = agents_api_config.transform_list_versions_request(
|
||||
name=name,
|
||||
api_base=litellm_params.get("api_base"),
|
||||
litellm_params=dict(litellm_params),
|
||||
)
|
||||
logging_obj.pre_call(
|
||||
input=name,
|
||||
api_key="",
|
||||
additional_args={"api_base": url, "headers": headers},
|
||||
)
|
||||
try:
|
||||
response = await async_httpx_client.get(
|
||||
url=url, headers=headers, params=params
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=agents_api_config)
|
||||
|
||||
logging_obj.post_call(original_response=response.text, additional_args={})
|
||||
return agents_api_config.transform_list_versions_response(
|
||||
raw_response=response, name=name
|
||||
)
|
||||
|
||||
|
||||
agents_http_handler = AgentsHTTPHandler()
|
||||
522
litellm/interactions/agents/main.py
Normal file
522
litellm/interactions/agents/main.py
Normal file
|
|
@ -0,0 +1,522 @@
|
|||
"""
|
||||
LiteLLM Agents API - Main Module
|
||||
|
||||
Usage:
|
||||
import litellm
|
||||
|
||||
# Create
|
||||
response = litellm.interactions.agents.create(
|
||||
name="waverunner",
|
||||
custom_llm_provider="gemini",
|
||||
api_key="...",
|
||||
base_agent="gemini-2.5-flash",
|
||||
instructions="You are a helpful assistant.",
|
||||
)
|
||||
|
||||
# List
|
||||
response = litellm.interactions.agents.list(api_key="...", custom_llm_provider="gemini")
|
||||
|
||||
# Get
|
||||
response = litellm.interactions.agents.get(name="waverunner", api_key="...")
|
||||
|
||||
# Delete
|
||||
result = litellm.interactions.agents.delete(name="waverunner", api_key="...")
|
||||
|
||||
# List versions
|
||||
result = litellm.interactions.agents.list_versions(name="waverunner", api_key="...")
|
||||
|
||||
# Async versions: acreate, alist, aget, adelete, alist_versions
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import contextvars
|
||||
from functools import partial
|
||||
from typing import Any, Coroutine, Dict, Optional, Union
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.interactions.agents.http_handler import agents_http_handler
|
||||
from litellm.interactions.agents.utils import get_provider_agents_api_config
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.agents import (
|
||||
AgentCreateResponse,
|
||||
AgentDeleteResult,
|
||||
AgentListResponse,
|
||||
AgentVersionsResponse,
|
||||
)
|
||||
from litellm.types.interactions import InteractionEnvironment
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.utils import client
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# Shared helpers #
|
||||
# ------------------------------------------------------------------ #
|
||||
|
||||
|
||||
def _get_agents_api_config(custom_llm_provider: str):
|
||||
config = get_provider_agents_api_config(custom_llm_provider)
|
||||
if config is None:
|
||||
raise litellm.BadRequestError(
|
||||
message=(
|
||||
f"Provider '{custom_llm_provider}' does not have a native "
|
||||
"agents API. Use the proxy POST /v1/agents endpoint to store "
|
||||
"agents locally."
|
||||
),
|
||||
model="",
|
||||
llm_provider=custom_llm_provider,
|
||||
)
|
||||
return config
|
||||
|
||||
|
||||
def _make_logging_obj(
|
||||
kwargs: Dict[str, Any],
|
||||
model: str,
|
||||
custom_llm_provider: str,
|
||||
call_type: str,
|
||||
optional_params: Dict[str, Any],
|
||||
) -> LiteLLMLoggingObj:
|
||||
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
|
||||
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
|
||||
litellm_logging_obj.update_from_kwargs(
|
||||
kwargs=kwargs,
|
||||
model=model,
|
||||
optional_params=optional_params,
|
||||
litellm_params={"litellm_call_id": litellm_call_id},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
return litellm_logging_obj
|
||||
|
||||
|
||||
# ================================================================== #
|
||||
# CREATE #
|
||||
# ================================================================== #
|
||||
|
||||
|
||||
@client
|
||||
async def acreate(
|
||||
name: str,
|
||||
base_agent: Optional[str] = None,
|
||||
instructions: Optional[str] = None,
|
||||
base_environment: Optional[InteractionEnvironment] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
**kwargs,
|
||||
) -> AgentCreateResponse:
|
||||
"""Async: Create a managed agent on the provider side."""
|
||||
local_vars = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
kwargs["acreate_agent"] = True
|
||||
func = partial(
|
||||
create,
|
||||
name=name,
|
||||
base_agent=base_agent,
|
||||
instructions=instructions,
|
||||
base_environment=base_environment,
|
||||
custom_llm_provider=custom_llm_provider or "gemini",
|
||||
extra_headers=extra_headers,
|
||||
extra_body=extra_body,
|
||||
timeout=timeout,
|
||||
**kwargs,
|
||||
)
|
||||
ctx = contextvars.copy_context()
|
||||
init_response = await loop.run_in_executor(None, partial(ctx.run, func))
|
||||
if asyncio.iscoroutine(init_response):
|
||||
return await init_response
|
||||
return init_response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=name,
|
||||
custom_llm_provider=custom_llm_provider or "gemini",
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
def create(
|
||||
name: str,
|
||||
base_agent: Optional[str] = None,
|
||||
instructions: Optional[str] = None,
|
||||
base_environment: Optional[InteractionEnvironment] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
**kwargs,
|
||||
) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]:
|
||||
"""
|
||||
Sync: Create a managed agent on the provider side.
|
||||
|
||||
Args:
|
||||
name: Name for the agent (required).
|
||||
base_agent: Base agent to derive from (e.g. "waverunner").
|
||||
instructions: System instructions for the agent.
|
||||
base_environment: Environment to fork from — an env_id string or a
|
||||
dict like ``{"type": "remote", "sources": [...]}``.
|
||||
custom_llm_provider: Provider to use, e.g. "gemini".
|
||||
extra_headers: Additional HTTP headers.
|
||||
extra_body: Additional request body fields.
|
||||
timeout: Request timeout.
|
||||
**kwargs: Forwarded to GenericLiteLLMParams (api_key, api_base, etc.).
|
||||
"""
|
||||
local_vars = locals()
|
||||
custom_llm_provider = (
|
||||
custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini"
|
||||
)
|
||||
try:
|
||||
_is_async = kwargs.pop("acreate_agent", False) is True
|
||||
if base_agent is not None:
|
||||
kwargs["base_agent"] = base_agent
|
||||
if instructions is not None:
|
||||
kwargs["instructions"] = instructions
|
||||
if base_environment is not None:
|
||||
kwargs["base_environment"] = base_environment
|
||||
kwargs.setdefault("custom_llm_provider", custom_llm_provider)
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
logging_obj = _make_logging_obj(
|
||||
kwargs, name, custom_llm_provider, "create_agent", {}
|
||||
)
|
||||
config = _get_agents_api_config(custom_llm_provider)
|
||||
return agents_http_handler.create_agent(
|
||||
agents_api_config=config,
|
||||
name=name,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
extra_body=extra_body,
|
||||
timeout=timeout,
|
||||
_is_async=_is_async,
|
||||
)
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=name,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
# ================================================================== #
|
||||
# LIST #
|
||||
# ================================================================== #
|
||||
|
||||
|
||||
@client
|
||||
async def alist(
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
**kwargs,
|
||||
) -> AgentListResponse:
|
||||
"""Async: List all agents on the provider side."""
|
||||
local_vars = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
kwargs["alist_agents"] = True
|
||||
func = partial(
|
||||
list,
|
||||
custom_llm_provider=custom_llm_provider or "gemini",
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout,
|
||||
**kwargs,
|
||||
)
|
||||
ctx = contextvars.copy_context()
|
||||
init_response = await loop.run_in_executor(None, partial(ctx.run, func))
|
||||
if asyncio.iscoroutine(init_response):
|
||||
return await init_response
|
||||
return init_response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model="",
|
||||
custom_llm_provider=custom_llm_provider or "gemini",
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
def list(
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
**kwargs,
|
||||
) -> Union[AgentListResponse, Coroutine[Any, Any, AgentListResponse]]:
|
||||
"""Sync: List all agents on the provider side."""
|
||||
local_vars = locals()
|
||||
custom_llm_provider = (
|
||||
custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini"
|
||||
)
|
||||
try:
|
||||
_is_async = kwargs.pop("alist_agents", False) is True
|
||||
kwargs.setdefault("custom_llm_provider", custom_llm_provider)
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
logging_obj = _make_logging_obj(
|
||||
kwargs, "", custom_llm_provider, "list_agents", {}
|
||||
)
|
||||
config = _get_agents_api_config(custom_llm_provider)
|
||||
return agents_http_handler.list_agents(
|
||||
agents_api_config=config,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout,
|
||||
_is_async=_is_async,
|
||||
)
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model="",
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
# ================================================================== #
|
||||
# GET #
|
||||
# ================================================================== #
|
||||
|
||||
|
||||
@client
|
||||
async def aget(
|
||||
name: str,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
**kwargs,
|
||||
) -> AgentCreateResponse:
|
||||
"""Async: Get a specific agent by name."""
|
||||
local_vars = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
kwargs["aget_agent"] = True
|
||||
func = partial(
|
||||
get,
|
||||
name=name,
|
||||
custom_llm_provider=custom_llm_provider or "gemini",
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout,
|
||||
**kwargs,
|
||||
)
|
||||
ctx = contextvars.copy_context()
|
||||
init_response = await loop.run_in_executor(None, partial(ctx.run, func))
|
||||
if asyncio.iscoroutine(init_response):
|
||||
return await init_response
|
||||
return init_response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=name,
|
||||
custom_llm_provider=custom_llm_provider or "gemini",
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
def get(
|
||||
name: str,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
**kwargs,
|
||||
) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]:
|
||||
"""Sync: Get a specific agent by name."""
|
||||
local_vars = locals()
|
||||
custom_llm_provider = (
|
||||
custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini"
|
||||
)
|
||||
try:
|
||||
_is_async = kwargs.pop("aget_agent", False) is True
|
||||
kwargs.setdefault("custom_llm_provider", custom_llm_provider)
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
logging_obj = _make_logging_obj(
|
||||
kwargs, name, custom_llm_provider, "get_agent", {"name": name}
|
||||
)
|
||||
config = _get_agents_api_config(custom_llm_provider)
|
||||
return agents_http_handler.get_agent(
|
||||
agents_api_config=config,
|
||||
name=name,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout,
|
||||
_is_async=_is_async,
|
||||
)
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=name,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
# ================================================================== #
|
||||
# DELETE #
|
||||
# ================================================================== #
|
||||
|
||||
|
||||
@client
|
||||
async def adelete(
|
||||
name: str,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
**kwargs,
|
||||
) -> AgentDeleteResult:
|
||||
"""Async: Delete a specific agent by name."""
|
||||
local_vars = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
kwargs["adelete_agent"] = True
|
||||
func = partial(
|
||||
delete,
|
||||
name=name,
|
||||
custom_llm_provider=custom_llm_provider or "gemini",
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout,
|
||||
**kwargs,
|
||||
)
|
||||
ctx = contextvars.copy_context()
|
||||
init_response = await loop.run_in_executor(None, partial(ctx.run, func))
|
||||
if asyncio.iscoroutine(init_response):
|
||||
return await init_response
|
||||
return init_response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=name,
|
||||
custom_llm_provider=custom_llm_provider or "gemini",
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
def delete(
|
||||
name: str,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
**kwargs,
|
||||
) -> Union[AgentDeleteResult, Coroutine[Any, Any, AgentDeleteResult]]:
|
||||
"""Sync: Delete a specific agent by name."""
|
||||
local_vars = locals()
|
||||
custom_llm_provider = (
|
||||
custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini"
|
||||
)
|
||||
try:
|
||||
_is_async = kwargs.pop("adelete_agent", False) is True
|
||||
kwargs.setdefault("custom_llm_provider", custom_llm_provider)
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
logging_obj = _make_logging_obj(
|
||||
kwargs, name, custom_llm_provider, "delete_agent", {"name": name}
|
||||
)
|
||||
config = _get_agents_api_config(custom_llm_provider)
|
||||
return agents_http_handler.delete_agent(
|
||||
agents_api_config=config,
|
||||
name=name,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout,
|
||||
_is_async=_is_async,
|
||||
)
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=name,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
# ================================================================== #
|
||||
# LIST VERSIONS #
|
||||
# ================================================================== #
|
||||
|
||||
|
||||
@client
|
||||
async def alist_versions(
|
||||
name: str,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
**kwargs,
|
||||
) -> AgentVersionsResponse:
|
||||
"""Async: List versions of a specific agent."""
|
||||
local_vars = locals()
|
||||
try:
|
||||
loop = asyncio.get_event_loop()
|
||||
kwargs["alist_agent_versions"] = True
|
||||
func = partial(
|
||||
list_versions,
|
||||
name=name,
|
||||
custom_llm_provider=custom_llm_provider or "gemini",
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout,
|
||||
**kwargs,
|
||||
)
|
||||
ctx = contextvars.copy_context()
|
||||
init_response = await loop.run_in_executor(None, partial(ctx.run, func))
|
||||
if asyncio.iscoroutine(init_response):
|
||||
return await init_response
|
||||
return init_response
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=name,
|
||||
custom_llm_provider=custom_llm_provider or "gemini",
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
def list_versions(
|
||||
name: str,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
extra_headers: Optional[Dict[str, Any]] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
**kwargs,
|
||||
) -> Union[AgentVersionsResponse, Coroutine[Any, Any, AgentVersionsResponse]]:
|
||||
"""Sync: List versions of a specific agent."""
|
||||
local_vars = locals()
|
||||
custom_llm_provider = (
|
||||
custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini"
|
||||
)
|
||||
try:
|
||||
_is_async = kwargs.pop("alist_agent_versions", False) is True
|
||||
kwargs.setdefault("custom_llm_provider", custom_llm_provider)
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
logging_obj = _make_logging_obj(
|
||||
kwargs, name, custom_llm_provider, "list_agent_versions", {"name": name}
|
||||
)
|
||||
config = _get_agents_api_config(custom_llm_provider)
|
||||
return agents_http_handler.list_agent_versions(
|
||||
agents_api_config=config,
|
||||
name=name,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout,
|
||||
_is_async=_is_async,
|
||||
)
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=name,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
original_exception=e,
|
||||
completion_kwargs=local_vars,
|
||||
extra_kwargs=kwargs,
|
||||
)
|
||||
23
litellm/interactions/agents/utils.py
Normal file
23
litellm/interactions/agents/utils.py
Normal file
|
|
@ -0,0 +1,23 @@
|
|||
"""
|
||||
Utility functions for the Agents API SDK.
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from litellm.llms.base_llm.agents.transformation import BaseAgentsAPIConfig
|
||||
|
||||
|
||||
def get_provider_agents_api_config(
|
||||
custom_llm_provider: Optional[str],
|
||||
) -> Optional[BaseAgentsAPIConfig]:
|
||||
"""
|
||||
Return a provider-specific BaseAgentsAPIConfig if the provider has a
|
||||
native agent-creation API, or None otherwise.
|
||||
"""
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
if custom_llm_provider == LlmProviders.GEMINI.value:
|
||||
from litellm.llms.gemini.agents.transformation import GeminiAgentsConfig
|
||||
|
||||
return GeminiAgentsConfig()
|
||||
return None
|
||||
|
|
@ -41,27 +41,55 @@ from litellm.types.interactions import (
|
|||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
|
||||
class InteractionsHTTPHandler:
|
||||
class _BaseHTTPHandler:
|
||||
"""
|
||||
Shared HTTP infrastructure for LiteLLM handler classes.
|
||||
|
||||
Provides common client resolution and error-mapping helpers so that
|
||||
handler subclasses (InteractionsHTTPHandler, AgentsHTTPHandler, …) do
|
||||
not duplicate this boilerplate.
|
||||
"""
|
||||
|
||||
def _handle_error(self, e: Exception, provider_config: Any) -> Exception:
|
||||
if isinstance(e, httpx.HTTPStatusError):
|
||||
return provider_config.get_error_class(
|
||||
error_message=e.response.text,
|
||||
status_code=e.response.status_code,
|
||||
headers=dict(e.response.headers),
|
||||
)
|
||||
return e
|
||||
|
||||
def _sync_client(
|
||||
self,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
client: Optional[HTTPHandler],
|
||||
) -> HTTPHandler:
|
||||
return client or _get_httpx_client(
|
||||
params={"ssl_verify": litellm_params.get("ssl_verify", None)}
|
||||
)
|
||||
|
||||
def _async_client(
|
||||
self,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
client: Optional[AsyncHTTPHandler],
|
||||
) -> AsyncHTTPHandler:
|
||||
# GenericLiteLLMParams.get uses getattr; an unset field is None, not the default.
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider") or "gemini"
|
||||
return client or get_async_httpx_client(
|
||||
llm_provider=litellm.LlmProviders(custom_llm_provider),
|
||||
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
|
||||
)
|
||||
|
||||
|
||||
class InteractionsHTTPHandler(_BaseHTTPHandler):
|
||||
"""
|
||||
HTTP handler for Interactions API requests.
|
||||
"""
|
||||
|
||||
def _handle_error(
|
||||
self,
|
||||
e: Exception,
|
||||
provider_config: BaseInteractionsAPIConfig,
|
||||
) -> Exception:
|
||||
"""Handle errors from HTTP requests."""
|
||||
if isinstance(e, httpx.HTTPStatusError):
|
||||
error_message = e.response.text
|
||||
status_code = e.response.status_code
|
||||
headers = dict(e.response.headers)
|
||||
return provider_config.get_error_class(
|
||||
error_message=error_message,
|
||||
status_code=status_code,
|
||||
headers=headers,
|
||||
)
|
||||
return e
|
||||
# _handle_error is inherited from _BaseHTTPHandler (accepts Any provider_config).
|
||||
# AgentsHTTPHandler also extends this class and passes BaseAgentsAPIConfig, which
|
||||
# is structurally compatible but a different type — keeping the override here with
|
||||
# BaseInteractionsAPIConfig would cause type errors in the subclass.
|
||||
|
||||
# =========================================================
|
||||
# CREATE INTERACTION
|
||||
|
|
|
|||
|
|
@ -2,7 +2,17 @@
|
|||
Streaming iterator for transforming Responses API stream to Interactions API stream.
|
||||
"""
|
||||
|
||||
from typing import Any, AsyncIterator, Dict, Iterator, Optional, cast
|
||||
from collections import deque
|
||||
from typing import (
|
||||
Any,
|
||||
AsyncIterator,
|
||||
Deque,
|
||||
Dict,
|
||||
Iterator,
|
||||
List,
|
||||
Optional,
|
||||
cast,
|
||||
)
|
||||
|
||||
from litellm.responses.streaming_iterator import (
|
||||
BaseResponsesAPIStreamingIterator,
|
||||
|
|
@ -29,7 +39,13 @@ class LiteLLMResponsesInteractionsStreamingIterator:
|
|||
|
||||
This class handles both sync and async iteration, transforming Responses API
|
||||
streaming events (output.text.delta, response.completed, etc.) to Interactions
|
||||
API streaming events (content.delta, interaction.complete, etc.).
|
||||
API streaming events.
|
||||
|
||||
Schema selection:
|
||||
- New schema (default, use_legacy_interactions_schema=False):
|
||||
interaction.created -> step.start -> step.delta ... -> step.stop -> interaction.completed
|
||||
- Legacy schema (use_legacy_interactions_schema=True, remove after June 8 2026):
|
||||
interaction.start -> content.start -> content.delta ... -> content.stop -> interaction.complete
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
|
|
@ -41,6 +57,8 @@ class LiteLLMResponsesInteractionsStreamingIterator:
|
|||
custom_llm_provider: Optional[str] = None,
|
||||
litellm_metadata: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
import litellm
|
||||
|
||||
self.model = model
|
||||
self.responses_stream_iterator = litellm_custom_stream_wrapper
|
||||
self.request_input = request_input
|
||||
|
|
@ -51,66 +69,156 @@ class LiteLLMResponsesInteractionsStreamingIterator:
|
|||
self.collected_text = ""
|
||||
self.sent_interaction_start = False
|
||||
self.sent_content_start = False
|
||||
# Capture the schema flag once at construction time so all events
|
||||
# emitted by this stream use a consistent schema, even if the global
|
||||
# flag is mutated mid-stream (e.g. by a config reload).
|
||||
self._use_legacy: bool = litellm.use_legacy_interactions_schema
|
||||
# Buffer of events that have been derived from upstream chunks but not
|
||||
# yet returned to the caller. A single Responses API chunk may expand
|
||||
# into multiple Interactions API events (e.g. the first text delta
|
||||
# produces interaction.created + step.start + step.delta), and the
|
||||
# terminal sequence on stream end may also span multiple events
|
||||
# (step.stop + interaction.completed).
|
||||
self._pending_events: Deque[InteractionsAPIStreamingResponse] = deque()
|
||||
# Tracks whether we've already emitted a terminal completion event so
|
||||
# the StopIteration fallback path doesn't double-emit.
|
||||
self._sent_completion_event = False
|
||||
# ID resolved from the first upstream chunk (item_id on a text delta or
|
||||
# response.id on response.created). Persisted so the EOF terminal
|
||||
# events stay correlated with the start events delivered earlier.
|
||||
self._interaction_id: Optional[str] = None
|
||||
|
||||
def _transform_responses_chunk_to_interactions_chunk(
|
||||
self,
|
||||
responses_chunk: ResponsesAPIStreamingResponse,
|
||||
) -> Optional[InteractionsAPIStreamingResponse]:
|
||||
# ------------------------------------------------------------------
|
||||
# Event builders
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _build_interaction_start_event(
|
||||
self, interaction_id: str
|
||||
) -> InteractionsAPIStreamingResponse:
|
||||
event_type = "interaction.start" if self._use_legacy else "interaction.created"
|
||||
return InteractionsAPIStreamingResponse(
|
||||
event_type=event_type,
|
||||
id=interaction_id,
|
||||
object="interaction",
|
||||
status="in_progress",
|
||||
model=self.model,
|
||||
)
|
||||
|
||||
def _build_content_start_event(
|
||||
self, interaction_id: str
|
||||
) -> InteractionsAPIStreamingResponse:
|
||||
if self._use_legacy:
|
||||
return InteractionsAPIStreamingResponse(
|
||||
event_type="content.start",
|
||||
id=interaction_id,
|
||||
object="content",
|
||||
delta={"type": "text", "text": ""},
|
||||
)
|
||||
return InteractionsAPIStreamingResponse(
|
||||
event_type="step.start",
|
||||
index=0,
|
||||
step={"type": "model_output", "content": []},
|
||||
)
|
||||
|
||||
def _build_text_delta_event(
|
||||
self, interaction_id: str, delta_text: str
|
||||
) -> InteractionsAPIStreamingResponse:
|
||||
if self._use_legacy:
|
||||
return InteractionsAPIStreamingResponse(
|
||||
event_type="content.delta",
|
||||
id=interaction_id,
|
||||
object="content",
|
||||
delta={"type": "text", "text": delta_text},
|
||||
)
|
||||
return InteractionsAPIStreamingResponse(
|
||||
event_type="step.delta",
|
||||
index=0,
|
||||
delta={"type": "text", "text": delta_text},
|
||||
)
|
||||
|
||||
def _build_content_stop_event(
|
||||
self, interaction_id: Optional[str]
|
||||
) -> InteractionsAPIStreamingResponse:
|
||||
if self._use_legacy:
|
||||
return InteractionsAPIStreamingResponse(
|
||||
event_type="content.stop",
|
||||
id=interaction_id,
|
||||
object="content",
|
||||
delta={"type": "text", "text": self.collected_text},
|
||||
)
|
||||
return InteractionsAPIStreamingResponse(
|
||||
event_type="step.stop",
|
||||
index=0,
|
||||
)
|
||||
|
||||
def _build_completion_event(
|
||||
self, response_id: str
|
||||
) -> InteractionsAPIStreamingResponse:
|
||||
if self._use_legacy:
|
||||
return InteractionsAPIStreamingResponse(
|
||||
event_type="interaction.complete",
|
||||
id=response_id,
|
||||
object="interaction",
|
||||
status="completed",
|
||||
model=self.model,
|
||||
outputs=[{"type": "text", "text": self.collected_text}],
|
||||
)
|
||||
return InteractionsAPIStreamingResponse(
|
||||
event_type="interaction.completed",
|
||||
id=response_id,
|
||||
object="interaction",
|
||||
status="completed",
|
||||
model=self.model,
|
||||
steps=[
|
||||
{
|
||||
"type": "model_output",
|
||||
"content": [{"type": "text", "text": self.collected_text}],
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Per-chunk transform (returns a list of events to enqueue)
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _events_for_chunk(
|
||||
self, responses_chunk: ResponsesAPIStreamingResponse
|
||||
) -> List[InteractionsAPIStreamingResponse]:
|
||||
"""
|
||||
Transform a Responses API streaming chunk to an Interactions API streaming chunk.
|
||||
Translate a single upstream Responses API chunk into the list of
|
||||
Interactions API events it should produce.
|
||||
|
||||
Responses API events:
|
||||
- output.text.delta -> content.delta
|
||||
- response.completed -> interaction.complete
|
||||
|
||||
Interactions API events:
|
||||
- interaction.start
|
||||
- content.start
|
||||
- content.delta
|
||||
- content.stop
|
||||
- interaction.complete
|
||||
Returning a list (rather than a single event) lets a chunk emit any
|
||||
synthetic start events that haven't been sent yet *together with* the
|
||||
actual delta event, so we never silently drop the chunk's payload.
|
||||
"""
|
||||
if not responses_chunk:
|
||||
return None
|
||||
return []
|
||||
|
||||
# Handle OutputTextDeltaEvent -> content.delta
|
||||
# Text delta: emit any missing start events, then the delta itself.
|
||||
if isinstance(responses_chunk, OutputTextDeltaEvent):
|
||||
delta_text = (
|
||||
responses_chunk.delta if isinstance(responses_chunk.delta, str) else ""
|
||||
)
|
||||
self.collected_text += delta_text
|
||||
interaction_id = (
|
||||
getattr(responses_chunk, "item_id", None) or f"interaction_{id(self)}"
|
||||
)
|
||||
if self._interaction_id is None:
|
||||
self._interaction_id = interaction_id
|
||||
|
||||
# Send interaction.start if not sent
|
||||
events: List[InteractionsAPIStreamingResponse] = []
|
||||
if not self.sent_interaction_start:
|
||||
self.sent_interaction_start = True
|
||||
return InteractionsAPIStreamingResponse(
|
||||
event_type="interaction.start",
|
||||
id=getattr(responses_chunk, "item_id", None)
|
||||
or f"interaction_{id(self)}",
|
||||
object="interaction",
|
||||
status="in_progress",
|
||||
model=self.model,
|
||||
)
|
||||
|
||||
# Send content.start if not sent
|
||||
events.append(self._build_interaction_start_event(interaction_id))
|
||||
if not self.sent_content_start:
|
||||
self.sent_content_start = True
|
||||
return InteractionsAPIStreamingResponse(
|
||||
event_type="content.start",
|
||||
id=getattr(responses_chunk, "item_id", None),
|
||||
object="content",
|
||||
delta={"type": "text", "text": ""},
|
||||
)
|
||||
events.append(self._build_content_start_event(interaction_id))
|
||||
events.append(self._build_text_delta_event(interaction_id, delta_text))
|
||||
return events
|
||||
|
||||
# Send content.delta
|
||||
return InteractionsAPIStreamingResponse(
|
||||
event_type="content.delta",
|
||||
id=getattr(responses_chunk, "item_id", None),
|
||||
object="content",
|
||||
delta={"text": delta_text},
|
||||
)
|
||||
|
||||
# Handle ResponseCreatedEvent or ResponseInProgressEvent -> interaction.start
|
||||
# Response created / in-progress: synthesize interaction start if we
|
||||
# haven't already sent one.
|
||||
if isinstance(responses_chunk, (ResponseCreatedEvent, ResponseInProgressEvent)):
|
||||
if not self.sent_interaction_start:
|
||||
self.sent_interaction_start = True
|
||||
|
|
@ -118,169 +226,136 @@ class LiteLLMResponsesInteractionsStreamingIterator:
|
|||
getattr(responses_chunk.response, "id", None)
|
||||
if hasattr(responses_chunk, "response")
|
||||
else None
|
||||
)
|
||||
return InteractionsAPIStreamingResponse(
|
||||
event_type="interaction.start",
|
||||
id=response_id or f"interaction_{id(self)}",
|
||||
object="interaction",
|
||||
status="in_progress",
|
||||
model=self.model,
|
||||
)
|
||||
) or f"interaction_{id(self)}"
|
||||
if self._interaction_id is None:
|
||||
self._interaction_id = response_id
|
||||
return [self._build_interaction_start_event(response_id)]
|
||||
return []
|
||||
|
||||
# Handle ResponseCompletedEvent -> interaction.complete
|
||||
# Response completed: emit step.stop (if content was started) followed
|
||||
# by the terminal completion event. Prefer the interaction id already
|
||||
# established by earlier events so consumers can correlate the start
|
||||
# and completion events by id (response.id may differ from the item_id
|
||||
# used to derive the initial id when the stream starts directly with a
|
||||
# text delta).
|
||||
if isinstance(responses_chunk, ResponseCompletedEvent):
|
||||
self.finished = True
|
||||
response = responses_chunk.response
|
||||
|
||||
# Send content.stop first if content was started
|
||||
if self.sent_content_start:
|
||||
# Note: We'll send this in the iterator, not here
|
||||
pass
|
||||
|
||||
# Send interaction.complete
|
||||
return InteractionsAPIStreamingResponse(
|
||||
event_type="interaction.complete",
|
||||
id=getattr(response, "id", None) or f"interaction_{id(self)}",
|
||||
object="interaction",
|
||||
status="completed",
|
||||
model=self.model,
|
||||
outputs=[
|
||||
{
|
||||
"type": "text",
|
||||
"text": self.collected_text,
|
||||
}
|
||||
],
|
||||
response_id = (
|
||||
self._interaction_id
|
||||
or getattr(response, "id", None)
|
||||
or f"interaction_{id(self)}"
|
||||
)
|
||||
|
||||
# For other event types, return None (skip)
|
||||
return None
|
||||
terminal: List[InteractionsAPIStreamingResponse] = []
|
||||
if self.sent_content_start:
|
||||
terminal.append(self._build_content_stop_event(response_id))
|
||||
terminal.append(self._build_completion_event(response_id))
|
||||
self._sent_completion_event = True
|
||||
return terminal
|
||||
|
||||
return []
|
||||
|
||||
def _build_terminal_events_on_eof(
|
||||
self,
|
||||
) -> List[InteractionsAPIStreamingResponse]:
|
||||
"""
|
||||
Build the events to flush when the upstream stream ends without a
|
||||
ResponseCompletedEvent. Ensures consumers always observe a terminal
|
||||
interaction.completed/interaction.complete carrying the full text.
|
||||
"""
|
||||
if self._sent_completion_event:
|
||||
return []
|
||||
|
||||
fallback_id = self._interaction_id or f"interaction_{id(self)}"
|
||||
terminal: List[InteractionsAPIStreamingResponse] = []
|
||||
if self.sent_content_start:
|
||||
terminal.append(self._build_content_stop_event(fallback_id))
|
||||
if self.sent_interaction_start or self.collected_text:
|
||||
terminal.append(self._build_completion_event(fallback_id))
|
||||
self._sent_completion_event = True
|
||||
return terminal
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Iteration
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def __iter__(self) -> Iterator[InteractionsAPIStreamingResponse]:
|
||||
"""Sync iterator implementation."""
|
||||
return self
|
||||
|
||||
def __next__(self) -> InteractionsAPIStreamingResponse:
|
||||
"""Get next chunk in sync mode."""
|
||||
if self._pending_events:
|
||||
return self._pending_events.popleft()
|
||||
|
||||
if self.finished:
|
||||
raise StopIteration
|
||||
|
||||
# Check if we have a pending interaction.complete to send
|
||||
if hasattr(self, "_pending_interaction_complete"):
|
||||
pending: InteractionsAPIStreamingResponse = getattr(
|
||||
self, "_pending_interaction_complete"
|
||||
)
|
||||
delattr(self, "_pending_interaction_complete")
|
||||
return pending
|
||||
|
||||
# Use a loop instead of recursion to avoid stack overflow
|
||||
sync_iterator = cast(
|
||||
SyncResponsesAPIStreamingIterator, self.responses_stream_iterator
|
||||
)
|
||||
while True:
|
||||
try:
|
||||
# Get next chunk from responses API stream
|
||||
chunk = next(sync_iterator)
|
||||
|
||||
# Transform chunk (chunk is already a ResponsesAPIStreamingResponse)
|
||||
transformed = self._transform_responses_chunk_to_interactions_chunk(
|
||||
chunk
|
||||
)
|
||||
|
||||
if transformed:
|
||||
# If we finished and content was started, send content.stop before interaction.complete
|
||||
if (
|
||||
self.finished
|
||||
and self.sent_content_start
|
||||
and transformed.event_type == "interaction.complete"
|
||||
):
|
||||
# Send content.stop first
|
||||
content_stop = InteractionsAPIStreamingResponse(
|
||||
event_type="content.stop",
|
||||
id=transformed.id,
|
||||
object="content",
|
||||
delta={"type": "text", "text": self.collected_text},
|
||||
)
|
||||
# Store the interaction.complete to send next
|
||||
self._pending_interaction_complete = transformed
|
||||
return content_stop
|
||||
return transformed
|
||||
|
||||
# If no transformation, continue to next chunk (loop continues)
|
||||
|
||||
except StopIteration:
|
||||
self.finished = True
|
||||
self._pending_events.extend(self._build_terminal_events_on_eof())
|
||||
if self._pending_events:
|
||||
return self._pending_events.popleft()
|
||||
raise
|
||||
|
||||
# Send final events if needed
|
||||
if self.sent_content_start:
|
||||
return InteractionsAPIStreamingResponse(
|
||||
event_type="content.stop",
|
||||
object="content",
|
||||
delta={"type": "text", "text": self.collected_text},
|
||||
)
|
||||
|
||||
raise StopIteration
|
||||
events = self._events_for_chunk(chunk)
|
||||
if events:
|
||||
self._pending_events.extend(events)
|
||||
return self._pending_events.popleft()
|
||||
|
||||
def __aiter__(self) -> AsyncIterator[InteractionsAPIStreamingResponse]:
|
||||
"""Async iterator implementation."""
|
||||
return self
|
||||
|
||||
async def __anext__(self) -> InteractionsAPIStreamingResponse:
|
||||
"""Get next chunk in async mode."""
|
||||
if self._pending_events:
|
||||
return self._pending_events.popleft()
|
||||
|
||||
if self.finished:
|
||||
raise StopAsyncIteration
|
||||
|
||||
# Check if we have a pending interaction.complete to send
|
||||
if hasattr(self, "_pending_interaction_complete"):
|
||||
pending: InteractionsAPIStreamingResponse = getattr(
|
||||
self, "_pending_interaction_complete"
|
||||
)
|
||||
delattr(self, "_pending_interaction_complete")
|
||||
return pending
|
||||
|
||||
# Use a loop instead of recursion to avoid stack overflow
|
||||
async_iterator = cast(
|
||||
ResponsesAPIStreamingIterator, self.responses_stream_iterator
|
||||
)
|
||||
while True:
|
||||
try:
|
||||
# Get next chunk from responses API stream
|
||||
chunk = await async_iterator.__anext__()
|
||||
|
||||
# Transform chunk (chunk is already a ResponsesAPIStreamingResponse)
|
||||
transformed = self._transform_responses_chunk_to_interactions_chunk(
|
||||
chunk
|
||||
)
|
||||
|
||||
if transformed:
|
||||
# If we finished and content was started, send content.stop before interaction.complete
|
||||
if (
|
||||
self.finished
|
||||
and self.sent_content_start
|
||||
and transformed.event_type == "interaction.complete"
|
||||
):
|
||||
# Send content.stop first
|
||||
content_stop = InteractionsAPIStreamingResponse(
|
||||
event_type="content.stop",
|
||||
id=transformed.id,
|
||||
object="content",
|
||||
delta={"type": "text", "text": self.collected_text},
|
||||
)
|
||||
# Store the interaction.complete to send next
|
||||
self._pending_interaction_complete = transformed
|
||||
return content_stop
|
||||
return transformed
|
||||
|
||||
# If no transformation, continue to next chunk (loop continues)
|
||||
|
||||
except StopAsyncIteration:
|
||||
self.finished = True
|
||||
self._pending_events.extend(self._build_terminal_events_on_eof())
|
||||
if self._pending_events:
|
||||
return self._pending_events.popleft()
|
||||
raise
|
||||
|
||||
# Send final events if needed
|
||||
if self.sent_content_start:
|
||||
return InteractionsAPIStreamingResponse(
|
||||
event_type="content.stop",
|
||||
object="content",
|
||||
delta={"type": "text", "text": self.collected_text},
|
||||
)
|
||||
events = self._events_for_chunk(chunk)
|
||||
if events:
|
||||
self._pending_events.extend(events)
|
||||
return self._pending_events.popleft()
|
||||
|
||||
raise StopAsyncIteration
|
||||
# ------------------------------------------------------------------
|
||||
# Backwards-compatible single-chunk transform (used by tests and any
|
||||
# external callers that drove the iterator chunk-by-chunk pre-fix).
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _transform_responses_chunk_to_interactions_chunk(
|
||||
self,
|
||||
responses_chunk: ResponsesAPIStreamingResponse,
|
||||
) -> Optional[InteractionsAPIStreamingResponse]:
|
||||
"""
|
||||
Compatibility shim: returns the *first* event produced for this chunk
|
||||
and queues any remaining events on ``self._pending_events`` so they
|
||||
are surfaced on subsequent calls/iterations.
|
||||
|
||||
Prefer ``_events_for_chunk`` in new code.
|
||||
"""
|
||||
events = self._events_for_chunk(responses_chunk)
|
||||
if not events:
|
||||
return None
|
||||
first = events[0]
|
||||
if len(events) > 1:
|
||||
self._pending_events.extend(events[1:])
|
||||
return first
|
||||
|
|
|
|||
|
|
@ -226,29 +226,37 @@ class LiteLLMResponsesInteractionsConfig:
|
|||
- Map status
|
||||
- Extract usage
|
||||
"""
|
||||
# Extract text from outputs
|
||||
outputs = []
|
||||
# Extract text from outputs and build both `outputs` (legacy) and `steps` (new schema).
|
||||
outputs: List[Dict[str, Any]] = []
|
||||
steps: List[Dict[str, Any]] = []
|
||||
if hasattr(responses_response, "output") and responses_response.output:
|
||||
for output_item in responses_response.output:
|
||||
# Use getattr with None default to safely access content
|
||||
content = getattr(output_item, "content", None)
|
||||
if content is not None:
|
||||
content_items = content if isinstance(content, list) else [content]
|
||||
model_output_contents: List[Dict[str, Any]] = []
|
||||
for content_item in content_items:
|
||||
# Check if content_item has text attribute
|
||||
text = getattr(content_item, "text", None)
|
||||
if text is not None:
|
||||
outputs.append(
|
||||
{
|
||||
"type": "text",
|
||||
"text": text,
|
||||
}
|
||||
)
|
||||
# Use independent dict instances so mutations to one
|
||||
# of `outputs` / `steps` don't leak into the other.
|
||||
outputs.append({"type": "text", "text": text})
|
||||
model_output_contents.append({"type": "text", "text": text})
|
||||
elif (
|
||||
isinstance(content_item, dict)
|
||||
and content_item.get("type") == "text"
|
||||
):
|
||||
outputs.append(content_item)
|
||||
outputs.append({**content_item})
|
||||
model_output_contents.append({**content_item})
|
||||
if model_output_contents:
|
||||
steps.append(
|
||||
{
|
||||
"type": "model_output",
|
||||
"content": model_output_contents,
|
||||
}
|
||||
)
|
||||
|
||||
# Convert created_at to ISO string
|
||||
created_at = getattr(responses_response, "created_at", None)
|
||||
|
|
@ -270,12 +278,14 @@ class LiteLLMResponsesInteractionsConfig:
|
|||
else:
|
||||
interactions_status = status
|
||||
|
||||
# Build interactions response
|
||||
# Build interactions response — populate both `outputs` (legacy schema) and
|
||||
# `steps` (new schema) so callers work regardless of which schema they expect.
|
||||
interactions_response_dict: Dict[str, Any] = {
|
||||
"id": getattr(responses_response, "id", ""),
|
||||
"object": "interaction",
|
||||
"status": interactions_status,
|
||||
"outputs": outputs,
|
||||
"steps": steps,
|
||||
"model": model or getattr(responses_response, "model", ""),
|
||||
"created": created,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -8,25 +8,25 @@ Per OpenAPI spec (https://ai.google.dev/static/api/interactions.openapi.json):
|
|||
|
||||
Usage:
|
||||
import litellm
|
||||
|
||||
|
||||
# Create an interaction with a model
|
||||
response = litellm.interactions.create(
|
||||
model="gemini-2.5-flash",
|
||||
input="Hello, how are you?"
|
||||
)
|
||||
|
||||
|
||||
# Create an interaction with an agent
|
||||
response = litellm.interactions.create(
|
||||
agent="deep-research-pro-preview-12-2025",
|
||||
input="Research the current state of cancer research"
|
||||
)
|
||||
|
||||
|
||||
# Async version
|
||||
response = await litellm.interactions.acreate(...)
|
||||
|
||||
|
||||
# Get an interaction
|
||||
response = litellm.interactions.get(interaction_id="...")
|
||||
|
||||
|
||||
# Delete an interaction
|
||||
result = litellm.interactions.delete(interaction_id="...")
|
||||
"""
|
||||
|
|
@ -48,6 +48,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
|||
from litellm.types.interactions import (
|
||||
CancelInteractionResult,
|
||||
DeleteInteractionResult,
|
||||
InteractionEnvironment,
|
||||
InteractionInput,
|
||||
InteractionsAPIResponse,
|
||||
InteractionsAPIStreamingResponse,
|
||||
|
|
@ -80,6 +81,8 @@ async def acreate(
|
|||
store: Optional[bool] = None,
|
||||
# Background execution
|
||||
background: Optional[bool] = None,
|
||||
# Agent execution environment ("remote", env id, or remote config object)
|
||||
environment: Optional[InteractionEnvironment] = None,
|
||||
# Response format
|
||||
response_modalities: Optional[List[str]] = None,
|
||||
response_format: Optional[Dict[str, Any]] = None,
|
||||
|
|
@ -109,6 +112,10 @@ async def acreate(
|
|||
stream: Whether to stream the response
|
||||
store: Whether to store the response for later retrieval
|
||||
background: Whether to run in background
|
||||
environment: Agent execution environment — ``"remote"``, an existing env id
|
||||
string, or a config object such as
|
||||
``{"type": "remote", "sources": [...]}`` /
|
||||
``{"type": "remote", "network": {...}}``
|
||||
response_modalities: Requested response modalities (TEXT, IMAGE, AUDIO)
|
||||
response_format: JSON schema for response format
|
||||
response_mime_type: MIME type of the response
|
||||
|
|
@ -144,6 +151,7 @@ async def acreate(
|
|||
stream=stream,
|
||||
store=store,
|
||||
background=background,
|
||||
environment=environment,
|
||||
response_modalities=response_modalities,
|
||||
response_format=response_format,
|
||||
response_mime_type=response_mime_type,
|
||||
|
|
@ -194,6 +202,8 @@ def create(
|
|||
store: Optional[bool] = None,
|
||||
# Background execution
|
||||
background: Optional[bool] = None,
|
||||
# Agent execution environment ("remote", env id, or remote config object)
|
||||
environment: Optional[InteractionEnvironment] = None,
|
||||
# Response format
|
||||
response_modalities: Optional[List[str]] = None,
|
||||
response_format: Optional[Dict[str, Any]] = None,
|
||||
|
|
@ -231,6 +241,10 @@ def create(
|
|||
stream: Whether to stream the response
|
||||
store: Whether to store the response for later retrieval
|
||||
background: Whether to run in background
|
||||
environment: Agent execution environment — ``"remote"``, an existing env id
|
||||
string, or a config object such as
|
||||
``{"type": "remote", "sources": [...]}`` /
|
||||
``{"type": "remote", "network": {...}}``
|
||||
response_modalities: Requested response modalities (TEXT, IMAGE, AUDIO)
|
||||
response_format: JSON schema for response format
|
||||
response_mime_type: MIME type of the response
|
||||
|
|
@ -252,7 +266,14 @@ def create(
|
|||
|
||||
litellm_params = GenericLiteLLMParams(**kwargs)
|
||||
|
||||
if model:
|
||||
# Routing logic:
|
||||
# - agent provided (no model, or model accidentally set to agent name) → gemini
|
||||
# - model provided → resolve provider via get_llm_provider (normal routing)
|
||||
if agent and model == agent:
|
||||
model = None
|
||||
if agent and not model:
|
||||
custom_llm_provider = custom_llm_provider or "gemini"
|
||||
elif model:
|
||||
model, custom_llm_provider, _, _ = litellm.get_llm_provider(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
|
|
|
|||
|
|
@ -101,10 +101,14 @@ class BaseInteractionsAPIStreamingIterator:
|
|||
)
|
||||
)
|
||||
|
||||
# Store the completed response (check for status=completed)
|
||||
if (
|
||||
streaming_response
|
||||
and getattr(streaming_response, "status", None) == "completed"
|
||||
# Store the completed response.
|
||||
# Legacy schema signals completion via status="completed".
|
||||
# New schema (Api-Revision: 2026-05-20) uses event_type="interaction.completed".
|
||||
# Remove the legacy check after June 8, 2026.
|
||||
if streaming_response and (
|
||||
getattr(streaming_response, "status", None) == "completed"
|
||||
or getattr(streaming_response, "event_type", None)
|
||||
== "interaction.completed"
|
||||
):
|
||||
self.completed_response = streaming_response
|
||||
self._handle_logging_completed_response()
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ INTERACTIONS_API_OPTIONAL_PARAMS = {
|
|||
"stream",
|
||||
"store",
|
||||
"background",
|
||||
"environment",
|
||||
"response_modalities",
|
||||
"response_format",
|
||||
"response_mime_type",
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ def get_supported_openai_params( # noqa: PLR0915
|
|||
request_type: Literal[
|
||||
"chat_completion", "embeddings", "transcription"
|
||||
] = "chat_completion",
|
||||
base_model: Optional[str] = None,
|
||||
) -> Optional[list]:
|
||||
"""
|
||||
Returns the supported openai params for a given model + provider
|
||||
|
|
@ -20,6 +21,11 @@ def get_supported_openai_params( # noqa: PLR0915
|
|||
get_supported_openai_params(model="anthropic.claude-3", custom_llm_provider="bedrock")
|
||||
```
|
||||
|
||||
Args:
|
||||
base_model: For Azure, the true underlying model (e.g. ``"azure/gpt-5.2"``)
|
||||
when the deployment name differs. Used for model-type detection so that
|
||||
non-standard deployment names route to the correct config.
|
||||
|
||||
Returns:
|
||||
- List if custom_llm_provider is mapped
|
||||
- None if unmapped
|
||||
|
|
@ -32,17 +38,21 @@ def get_supported_openai_params( # noqa: PLR0915
|
|||
|
||||
if custom_llm_provider in LlmProvidersSet:
|
||||
provider_config = litellm.ProviderConfigManager.get_provider_chat_config(
|
||||
model=model, provider=LlmProviders(custom_llm_provider)
|
||||
model=model,
|
||||
provider=LlmProviders(custom_llm_provider),
|
||||
base_model=base_model,
|
||||
)
|
||||
elif custom_llm_provider.split("/")[0] in LlmProvidersSet:
|
||||
provider_config = litellm.ProviderConfigManager.get_provider_chat_config(
|
||||
model=model, provider=LlmProviders(custom_llm_provider.split("/")[0])
|
||||
model=model,
|
||||
provider=LlmProviders(custom_llm_provider.split("/")[0]),
|
||||
base_model=base_model,
|
||||
)
|
||||
else:
|
||||
provider_config = None
|
||||
|
||||
if provider_config and request_type == "chat_completion":
|
||||
return provider_config.get_supported_openai_params(model=model)
|
||||
return provider_config.get_supported_openai_params(model=base_model or model)
|
||||
|
||||
if custom_llm_provider == "bedrock":
|
||||
return litellm.AmazonConverseConfig().get_supported_openai_params(model=model)
|
||||
|
|
@ -130,16 +140,23 @@ def get_supported_openai_params( # noqa: PLR0915
|
|||
model=model
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
if litellm.AzureOpenAIO1Config().is_o_series_model(model=model):
|
||||
_azure_detection_model = base_model or model
|
||||
if litellm.AzureOpenAIO1Config().is_o_series_model(
|
||||
model=_azure_detection_model
|
||||
):
|
||||
return litellm.AzureOpenAIO1Config().get_supported_openai_params(
|
||||
model=model
|
||||
model=_azure_detection_model
|
||||
)
|
||||
elif litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(model=model):
|
||||
elif litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(
|
||||
model=_azure_detection_model
|
||||
):
|
||||
return litellm.AzureOpenAIGPT5Config().get_supported_openai_params(
|
||||
model=model
|
||||
model=_azure_detection_model
|
||||
)
|
||||
else:
|
||||
return litellm.AzureOpenAIConfig().get_supported_openai_params(model=model)
|
||||
return litellm.AzureOpenAIConfig().get_supported_openai_params(
|
||||
model=_azure_detection_model
|
||||
)
|
||||
elif custom_llm_provider == "openrouter":
|
||||
return litellm.OpenrouterConfig().get_supported_openai_params(model=model)
|
||||
elif custom_llm_provider == "vercel_ai_gateway":
|
||||
|
|
|
|||
|
|
@ -994,10 +994,8 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
try:
|
||||
# [Non-blocking Extra Debug Information in metadata]
|
||||
if turn_off_message_logging is True:
|
||||
_metadata["raw_request"] = (
|
||||
"redacted by litellm. \
|
||||
_metadata["raw_request"] = "redacted by litellm. \
|
||||
'litellm.turn_off_message_logging=True'"
|
||||
)
|
||||
else:
|
||||
curl_command = self._get_request_curl_command(
|
||||
api_base=additional_args.get("api_base", ""),
|
||||
|
|
@ -1031,12 +1029,8 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
error=str(e),
|
||||
)
|
||||
)
|
||||
_metadata["raw_request"] = (
|
||||
"Unable to Log \
|
||||
raw request: {}".format(
|
||||
str(e)
|
||||
)
|
||||
)
|
||||
_metadata["raw_request"] = "Unable to Log \
|
||||
raw request: {}".format(str(e))
|
||||
if getattr(self, "logger_fn", None) and callable(self.logger_fn):
|
||||
try:
|
||||
self.logger_fn(
|
||||
|
|
@ -1769,9 +1763,12 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
self.model_call_details["response_cost"] = 0.0
|
||||
elif "response_cost" in hidden_params:
|
||||
self.model_call_details["response_cost"] = hidden_params["response_cost"]
|
||||
elif self.model_call_details.get("response_cost") is not None:
|
||||
elif (
|
||||
existing_cost := self.model_call_details.get("response_cost")
|
||||
) is not None and existing_cost != 0:
|
||||
# Preserve response_cost if already calculated (e.g., by pass-through
|
||||
# handlers like Gemini/Vertex which call completion_cost directly)
|
||||
# handlers like Gemini/Vertex which call completion_cost directly).
|
||||
# Do not preserve 0 from failure_handler on intermediate router retries.
|
||||
pass
|
||||
else:
|
||||
self.model_call_details["response_cost"] = self._response_cost_calculator(
|
||||
|
|
@ -5143,13 +5140,17 @@ class StandardLoggingPayloadSetup:
|
|||
) -> StandardLoggingPayloadErrorInformation:
|
||||
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
|
||||
|
||||
# Check for 'code' first (used by ProxyException), then fall back to 'status_code' (used by LiteLLM exceptions)
|
||||
# Ensure error_code is always a string for Prisma Python JSON field compatibility
|
||||
# ProxyException uses .code, LiteLLM exceptions use .status_code,
|
||||
# httpx.HTTPStatusError exposes status only as .response.status_code.
|
||||
# Stringified for Prisma JSON compatibility.
|
||||
error_code_attr = getattr(original_exception, "code", None)
|
||||
if error_code_attr is not None and str(error_code_attr) not in ("", "None"):
|
||||
error_status: str = str(error_code_attr)
|
||||
else:
|
||||
status_code_attr = getattr(original_exception, "status_code", None)
|
||||
if status_code_attr is None:
|
||||
response_attr = getattr(original_exception, "response", None)
|
||||
status_code_attr = getattr(response_attr, "status_code", None)
|
||||
error_status = str(status_code_attr) if status_code_attr is not None else ""
|
||||
error_class: str = (
|
||||
str(original_exception.__class__.__name__) if original_exception else ""
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ from typing import (
|
|||
cast,
|
||||
)
|
||||
|
||||
import litellm
|
||||
from litellm import verbose_logger
|
||||
from litellm.router_utils.batch_utils import InMemoryFile
|
||||
from litellm.types.llms.openai import (
|
||||
|
|
@ -1170,9 +1171,16 @@ def migrate_file_to_image_url(
|
|||
ChatCompletionImageUrlObject,
|
||||
)
|
||||
|
||||
file_id = message["file"].get("file_id")
|
||||
file_data = message["file"].get("file_data")
|
||||
format = message["file"].get("format")
|
||||
file_sub = message.get("file")
|
||||
if file_sub is None:
|
||||
raise litellm.BadRequestError(
|
||||
message="Content block has type='file' but is missing the required 'file' field",
|
||||
model=None,
|
||||
llm_provider=None,
|
||||
)
|
||||
file_id = file_sub.get("file_id")
|
||||
file_data = file_sub.get("file_data")
|
||||
format = file_sub.get("format")
|
||||
if not file_id and not file_data:
|
||||
raise ValueError("file_id and file_data are both None")
|
||||
image_url_object = ChatCompletionImageObject(
|
||||
|
|
@ -1196,12 +1204,8 @@ def get_last_user_message(messages: List[AllMessageValues]) -> Optional[str]:
|
|||
{"role": "assistant", "content": "I'm good, thank you!"},
|
||||
{"role": "user", "content": "What is the weather in Tokyo?"},
|
||||
]
|
||||
get_user_prompt(messages) -> "What is the weather in Tokyo?"
|
||||
get_last_user_message(messages) -> "What is the weather in Tokyo?"
|
||||
"""
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
convert_content_list_to_str,
|
||||
)
|
||||
|
||||
if not messages:
|
||||
return None
|
||||
|
||||
|
|
|
|||
|
|
@ -1233,6 +1233,7 @@ def infer_protocol_value(
|
|||
|
||||
def _gemini_tool_call_invoke_helper(
|
||||
function_call_params: ChatCompletionToolCallFunctionChunk,
|
||||
tool_call_id: Optional[str] = None,
|
||||
) -> Optional[VertexFunctionCall]:
|
||||
name = function_call_params.get("name", "") or ""
|
||||
arguments = function_call_params.get("arguments", "")
|
||||
|
|
@ -1248,6 +1249,10 @@ def _gemini_tool_call_invoke_helper(
|
|||
name=name,
|
||||
args=arguments_dict,
|
||||
)
|
||||
if tool_call_id:
|
||||
clean_id = tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1)[0]
|
||||
if clean_id:
|
||||
function_call["id"] = clean_id
|
||||
return function_call
|
||||
|
||||
|
||||
|
|
@ -1339,6 +1344,7 @@ def _get_dummy_thought_signature() -> str:
|
|||
def convert_to_gemini_tool_call_invoke(
|
||||
message: ChatCompletionAssistantMessage,
|
||||
model: Optional[str] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
) -> List[VertexPartType]:
|
||||
"""
|
||||
OpenAI tool invokes:
|
||||
|
|
@ -1384,12 +1390,26 @@ def convert_to_gemini_tool_call_invoke(
|
|||
tool_calls = message.get("tool_calls", None)
|
||||
function_call = message.get("function_call", None)
|
||||
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
VertexGeminiConfig,
|
||||
)
|
||||
|
||||
forward_tool_call_id = bool(
|
||||
model
|
||||
and VertexGeminiConfig._forward_gemini_function_call_id(
|
||||
model, custom_llm_provider
|
||||
)
|
||||
)
|
||||
|
||||
if tool_calls is not None:
|
||||
for idx, tool in enumerate(tool_calls):
|
||||
if "function" in tool:
|
||||
gemini_function_call: Optional[VertexFunctionCall] = (
|
||||
_gemini_tool_call_invoke_helper(
|
||||
function_call_params=tool["function"]
|
||||
function_call_params=tool["function"],
|
||||
tool_call_id=(
|
||||
tool.get("id") if forward_tool_call_id else None
|
||||
),
|
||||
)
|
||||
)
|
||||
if gemini_function_call is not None:
|
||||
|
|
@ -1429,10 +1449,6 @@ def convert_to_gemini_tool_call_invoke(
|
|||
thought_signature = provider_fields.get("thought_signature")
|
||||
|
||||
# If no signature found and model is gemini-3, use dummy signature
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
VertexGeminiConfig,
|
||||
)
|
||||
|
||||
if (
|
||||
not thought_signature
|
||||
and model
|
||||
|
|
@ -1462,6 +1478,8 @@ def convert_to_gemini_tool_call_invoke(
|
|||
def convert_to_gemini_tool_call_result( # noqa: PLR0915
|
||||
message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage],
|
||||
last_message_with_tool_calls: Optional[dict],
|
||||
model: Optional[str] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
) -> Union[VertexPartType, List[VertexPartType]]:
|
||||
"""
|
||||
OpenAI message with a tool result looks like:
|
||||
|
|
@ -1602,6 +1620,23 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915
|
|||
):
|
||||
name = tool.get("function", {}).get("name", "")
|
||||
|
||||
# Echo the OpenAI tool_call_id on functionResponse (strip thought-signature suffix).
|
||||
# Only Google AI Studio Gemini 3+ accepts `id` on function_response parts.
|
||||
# Vertex AI and older Gemini models reject the field with HTTP 400.
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
VertexGeminiConfig,
|
||||
)
|
||||
|
||||
gemini_call_id: Optional[str] = None
|
||||
if model and VertexGeminiConfig._forward_gemini_function_call_id(
|
||||
model, custom_llm_provider
|
||||
):
|
||||
raw_tool_call_id = message.get("tool_call_id")
|
||||
if raw_tool_call_id and isinstance(raw_tool_call_id, str):
|
||||
stripped_id = raw_tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1)[0]
|
||||
if stripped_id:
|
||||
gemini_call_id = stripped_id
|
||||
|
||||
if not name:
|
||||
raise Exception(
|
||||
"Missing corresponding tool call for tool response message. Received - message={}, last_message_with_tool_calls={}".format(
|
||||
|
|
@ -1632,6 +1667,8 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915
|
|||
name=name,
|
||||
response=response_data, # type: ignore
|
||||
)
|
||||
if gemini_call_id:
|
||||
_function_response["id"] = gemini_call_id
|
||||
|
||||
# Create part with function_response, and optionally inline_data for images (Computer Use)
|
||||
_part: VertexPartType = {"function_response": _function_response}
|
||||
|
|
@ -2057,9 +2094,16 @@ def anthropic_process_openai_file_message(
|
|||
AnthropicMessagesContainerUploadParam,
|
||||
]:
|
||||
file_message = cast(ChatCompletionFileObject, message)
|
||||
file_data = file_message["file"].get("file_data")
|
||||
file_id = file_message["file"].get("file_id")
|
||||
format = file_message["file"].get("format")
|
||||
file_sub = file_message.get("file")
|
||||
if file_sub is None:
|
||||
raise litellm.BadRequestError(
|
||||
message="Content block has type='file' but is missing the required 'file' field",
|
||||
model=None,
|
||||
llm_provider="anthropic",
|
||||
)
|
||||
file_data = file_sub.get("file_data")
|
||||
file_id = file_sub.get("file_id")
|
||||
format = file_sub.get("format")
|
||||
if file_data:
|
||||
image_chunk = convert_to_anthropic_image_obj(
|
||||
openai_image_url=file_data,
|
||||
|
|
@ -4879,7 +4923,13 @@ class BedrockConverseMessagesProcessor:
|
|||
|
||||
@staticmethod
|
||||
def _process_file_message(message: ChatCompletionFileObject) -> BedrockContentBlock:
|
||||
file_message = message["file"]
|
||||
file_message = message.get("file")
|
||||
if file_message is None:
|
||||
raise litellm.BadRequestError(
|
||||
message="Content block has type='file' but is missing the required 'file' field",
|
||||
model=None,
|
||||
llm_provider="bedrock",
|
||||
)
|
||||
file_data = file_message.get("file_data")
|
||||
file_id = file_message.get("file_id")
|
||||
|
||||
|
|
@ -4900,7 +4950,13 @@ class BedrockConverseMessagesProcessor:
|
|||
async def _async_process_file_message(
|
||||
message: ChatCompletionFileObject,
|
||||
) -> BedrockContentBlock:
|
||||
file_message = message["file"]
|
||||
file_message = message.get("file")
|
||||
if file_message is None:
|
||||
raise litellm.BadRequestError(
|
||||
message="Content block has type='file' but is missing the required 'file' field",
|
||||
model=None,
|
||||
llm_provider="bedrock",
|
||||
)
|
||||
file_data = file_message.get("file_data")
|
||||
file_id = file_message.get("file_id")
|
||||
format = file_message.get("format")
|
||||
|
|
@ -5534,9 +5590,7 @@ def default_response_schema_prompt(response_schema: dict) -> str:
|
|||
prompt_str = """Use this JSON schema:
|
||||
```json
|
||||
{}
|
||||
```""".format(
|
||||
response_schema
|
||||
)
|
||||
```""".format(response_schema)
|
||||
return prompt_str
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
"""
|
||||
This is a cache for LangfuseLoggers.
|
||||
|
||||
Langfuse Python SDK initializes a thread for each client.
|
||||
Langfuse Python SDK initializes a thread for each client.
|
||||
|
||||
This ensures we do
|
||||
This ensures we do
|
||||
1. Proper cleanup of Langfuse initialized clients.
|
||||
2. Re-use created langfuse clients.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1506,9 +1506,21 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
optional_params["metadata"] = {"user_id": value}
|
||||
elif param == "thinking":
|
||||
optional_params["thinking"] = value
|
||||
elif param == "reasoning_effort" and isinstance(value, str):
|
||||
elif param == "reasoning_effort":
|
||||
# Accept both string ("low") and dict ({"effort": "low",
|
||||
# "summary": "concise"}). The Responses->Chat parser keeps the
|
||||
# full dict when `summary` is set (see #25359), so a dict here
|
||||
# is the standard shape Otto/OpenAI-Responses-Bridge callers
|
||||
# send. Coerce to the effort string before mapping — same
|
||||
# shape-tolerance the GPT-5 path already implements in
|
||||
# `_normalize_reasoning_effort_for_chat_completion`.
|
||||
effort_value = value
|
||||
if isinstance(effort_value, dict):
|
||||
effort_value = effort_value.get("effort")
|
||||
if not isinstance(effort_value, str):
|
||||
continue
|
||||
mapped_thinking = AnthropicConfig._map_reasoning_effort(
|
||||
reasoning_effort=value,
|
||||
reasoning_effort=effort_value,
|
||||
model=model,
|
||||
llm_provider=self.custom_llm_provider or "anthropic",
|
||||
)
|
||||
|
|
@ -1519,12 +1531,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
optional_params["thinking"] = mapped_thinking
|
||||
if AnthropicConfig._is_adaptive_thinking_model(model):
|
||||
mapped_effort = REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT.get(
|
||||
value
|
||||
effort_value
|
||||
)
|
||||
if mapped_effort is None:
|
||||
AnthropicConfig._raise_invalid_reasoning_effort(
|
||||
model=model,
|
||||
value=value,
|
||||
value=effort_value,
|
||||
llm_provider=self.custom_llm_provider or "anthropic",
|
||||
)
|
||||
optional_params["output_config"] = {"effort": mapped_effort}
|
||||
|
|
|
|||
|
|
@ -1476,7 +1476,7 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
for choice in choices:
|
||||
if choice.delta.content is not None and len(choice.delta.content) > 0:
|
||||
text += choice.delta.content
|
||||
if choice.delta.tool_calls is not None:
|
||||
if choice.delta.tool_calls:
|
||||
partial_json = ""
|
||||
for tool in choice.delta.tool_calls:
|
||||
if (
|
||||
|
|
|
|||
|
|
@ -13,7 +13,6 @@ from typing import Any, AsyncIterator, Dict, List, Optional, cast
|
|||
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SSE parsing helpers (module-level to keep the class lean)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -293,6 +293,12 @@ async def anthropic_messages(
|
|||
api_base=api_base,
|
||||
client=client,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
# messages were already empty-text-block sanitized at the top of this
|
||||
# function and are NOT reassigned before this dispatch, so the handler
|
||||
# can skip its (otherwise redundant) second full-messages scan. Passed
|
||||
# explicitly (not via **kwargs) so it only affects this direct
|
||||
# dispatch -- interceptor / sync entry points still sanitize.
|
||||
_litellm_messages_presanitized=True,
|
||||
**kwargs,
|
||||
)
|
||||
ctx = contextvars.copy_context()
|
||||
|
|
@ -351,10 +357,14 @@ def anthropic_messages_handler(
|
|||
"""
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
# Sanitize empty text blocks here too so the sync entry point
|
||||
# Sanitize empty text blocks so the sync entry point
|
||||
# (litellm.messages.create -> anthropic_messages_handler) gets the same
|
||||
# protection as the async wrapper. Idempotent when called twice.
|
||||
messages = strip_empty_text_blocks_from_anthropic_messages(messages)
|
||||
# protection as the async wrapper. The async wrapper already sanitized and
|
||||
# does not reassign messages before dispatch, so it sets
|
||||
# ``_litellm_messages_presanitized`` to skip this redundant second
|
||||
# full-messages scan. Pop it so it never leaks into provider params.
|
||||
if not kwargs.pop("_litellm_messages_presanitized", False):
|
||||
messages = strip_empty_text_blocks_from_anthropic_messages(messages)
|
||||
|
||||
metadata = validate_anthropic_api_metadata(metadata)
|
||||
|
||||
|
|
|
|||
|
|
@ -312,7 +312,10 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
)
|
||||
|
||||
####### get required params for all anthropic messages requests ######
|
||||
verbose_logger.debug(f"TRANSFORMATION DEBUG - Messages: {messages}")
|
||||
# Lazy %s: the f-string previously stringified the entire messages
|
||||
# payload on every request regardless of log level (a full scan of the
|
||||
# request body on the hot path). Defer it to when DEBUG is enabled.
|
||||
verbose_logger.debug("TRANSFORMATION DEBUG - Messages: %s", messages)
|
||||
|
||||
# Auto-strip advisor blocks from history if advisor tool is absent.
|
||||
# Prevents Anthropic 400: advisor_tool_result in history requires advisor tool.
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
from typing import Any, Dict, List, cast, get_type_hints
|
||||
from functools import lru_cache
|
||||
from typing import Any, Dict, FrozenSet, List, cast, get_type_hints
|
||||
|
||||
from litellm.types.llms.anthropic import AnthropicMessagesRequestOptionalParams
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
|
|
@ -6,6 +7,18 @@ from litellm.types.llms.anthropic_messages.anthropic_response import (
|
|||
)
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _anthropic_messages_optional_param_keys() -> FrozenSet[str]:
|
||||
"""
|
||||
Valid AnthropicMessagesRequestOptionalParams keys.
|
||||
|
||||
``typing.get_type_hints`` is ~80us/call and this TypedDict is static, so
|
||||
resolving it once per process instead of once per request removes a fixed
|
||||
full-pass cost from the /v1/messages request-parse path.
|
||||
"""
|
||||
return frozenset(get_type_hints(AnthropicMessagesRequestOptionalParams).keys())
|
||||
|
||||
|
||||
class AnthropicMessagesRequestUtils:
|
||||
@staticmethod
|
||||
def get_requested_anthropic_messages_optional_param(
|
||||
|
|
@ -20,7 +33,7 @@ class AnthropicMessagesRequestUtils:
|
|||
Returns:
|
||||
AnthropicMessagesRequestOptionalParams instance with only the valid parameters
|
||||
"""
|
||||
valid_keys = get_type_hints(AnthropicMessagesRequestOptionalParams).keys()
|
||||
valid_keys = _anthropic_messages_optional_param_keys()
|
||||
filtered_params = {
|
||||
k: v for k, v in params.items() if k in valid_keys and v is not None
|
||||
}
|
||||
|
|
|
|||
3
litellm/llms/azure/audio_transcription/__init__.py
Normal file
3
litellm/llms/azure/audio_transcription/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from .transformation import AzureSpeechAudioTranscriptionConfig
|
||||
|
||||
__all__ = ["AzureSpeechAudioTranscriptionConfig"]
|
||||
224
litellm/llms/azure/audio_transcription/transformation.py
Normal file
224
litellm/llms/azure/audio_transcription/transformation.py
Normal file
|
|
@ -0,0 +1,224 @@
|
|||
"""
|
||||
Azure AI Speech (Cognitive Services) speech-to-text transformation.
|
||||
|
||||
Maps OpenAI-compatible audio transcription calls to Azure Speech REST
|
||||
recognition for short audio.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
from urllib.parse import urlencode, urlparse
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.litellm_core_utils.audio_utils.utils import process_audio_file
|
||||
from litellm.llms.base_llm.audio_transcription.transformation import (
|
||||
AudioTranscriptionRequestData,
|
||||
BaseAudioTranscriptionConfig,
|
||||
)
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
OpenAIAudioTranscriptionOptionalParams,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.utils import FileTypes, TranscriptionResponse
|
||||
|
||||
|
||||
class AzureSpeechAudioTranscriptionException(BaseLLMException):
|
||||
pass
|
||||
|
||||
|
||||
class AzureSpeechAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
|
||||
"""
|
||||
Configuration for Azure AI Speech (Cognitive Services) STT.
|
||||
|
||||
Reference:
|
||||
https://learn.microsoft.com/en-us/azure/ai-services/speech-service/rest-speech-to-text-short
|
||||
"""
|
||||
|
||||
COGNITIVE_SERVICES_DOMAIN = "api.cognitive.microsoft.com"
|
||||
STT_SPEECH_DOMAIN = "stt.speech.microsoft.com"
|
||||
STT_ENDPOINT_PATH = "/speech/recognition/conversation/cognitiveservices/v1"
|
||||
DEFAULT_LANGUAGE = "en-US"
|
||||
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> List[OpenAIAudioTranscriptionOptionalParams]:
|
||||
return ["language", "response_format"]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
supported_params = self.get_supported_openai_params(model=model)
|
||||
for key, value in non_default_params.items():
|
||||
if key in supported_params:
|
||||
optional_params[key] = value
|
||||
return optional_params
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
api_key = api_key or get_secret_str("AZURE_SPEECH_API_KEY")
|
||||
if not api_key:
|
||||
raise AzureSpeechAudioTranscriptionException(
|
||||
message="api_key is required for Azure AI Speech transcription.",
|
||||
status_code=401,
|
||||
)
|
||||
|
||||
validated_headers = headers.copy()
|
||||
validated_headers["Ocp-Apim-Subscription-Key"] = api_key
|
||||
validated_headers["Content-Type"] = validated_headers.get(
|
||||
"Content-Type", "audio/wav"
|
||||
)
|
||||
validated_headers["Accept"] = "application/json"
|
||||
return validated_headers
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
api_base = api_base or get_secret_str("AZURE_SPEECH_API_BASE")
|
||||
if api_base is None:
|
||||
raise AzureSpeechAudioTranscriptionException(
|
||||
message=(
|
||||
"api_base is required for Azure AI Speech transcription. "
|
||||
"Use a Cognitive Services endpoint like "
|
||||
"https://{region}.api.cognitive.microsoft.com or an STT "
|
||||
"endpoint like https://{region}.stt.speech.microsoft.com."
|
||||
),
|
||||
status_code=400,
|
||||
)
|
||||
|
||||
base_url = self._resolve_stt_base_url(api_base=api_base)
|
||||
query_params = {
|
||||
"language": optional_params.get("language", self.DEFAULT_LANGUAGE),
|
||||
"format": self._get_azure_response_format(
|
||||
optional_params.get("response_format")
|
||||
),
|
||||
}
|
||||
return f"{base_url}{self.STT_ENDPOINT_PATH}?{urlencode(query_params)}"
|
||||
|
||||
def transform_audio_transcription_request(
|
||||
self,
|
||||
model: str,
|
||||
audio_file: FileTypes,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
) -> AudioTranscriptionRequestData:
|
||||
processed_audio = process_audio_file(audio_file)
|
||||
return AudioTranscriptionRequestData(
|
||||
data=processed_audio.file_content,
|
||||
files=None,
|
||||
content_type=processed_audio.content_type,
|
||||
)
|
||||
|
||||
def transform_audio_transcription_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
) -> TranscriptionResponse:
|
||||
response_json = raw_response.json()
|
||||
recognition_status = response_json.get("RecognitionStatus")
|
||||
if recognition_status is not None and recognition_status != "Success":
|
||||
raise AzureSpeechAudioTranscriptionException(
|
||||
message=(
|
||||
"Azure AI Speech transcription failed with "
|
||||
f"RecognitionStatus={recognition_status}."
|
||||
),
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
|
||||
text = self._extract_text(response_json)
|
||||
response = TranscriptionResponse(text=text)
|
||||
response._hidden_params = response_json
|
||||
return response
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
) -> BaseLLMException:
|
||||
return AzureSpeechAudioTranscriptionException(
|
||||
message=error_message,
|
||||
status_code=status_code,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
def _resolve_stt_base_url(self, api_base: str) -> str:
|
||||
api_base = api_base.rstrip("/")
|
||||
parsed_url = urlparse(api_base)
|
||||
hostname = parsed_url.hostname or ""
|
||||
|
||||
if self._is_cognitive_services_endpoint(hostname=hostname):
|
||||
region = self._extract_region_from_hostname(
|
||||
hostname=hostname, domain=self.COGNITIVE_SERVICES_DOMAIN
|
||||
)
|
||||
return self._build_stt_base_url(region=region)
|
||||
|
||||
if self._is_stt_endpoint(hostname=hostname):
|
||||
return f"{parsed_url.scheme}://{hostname}"
|
||||
|
||||
if self._is_azure_openai_endpoint(hostname=hostname):
|
||||
raise AzureSpeechAudioTranscriptionException(
|
||||
message=(
|
||||
"Azure AI Speech transcription requires a Cognitive Services "
|
||||
"or STT Speech endpoint, not an Azure OpenAI endpoint."
|
||||
),
|
||||
status_code=400,
|
||||
)
|
||||
|
||||
return api_base
|
||||
|
||||
def _is_cognitive_services_endpoint(self, hostname: str) -> bool:
|
||||
return hostname == self.COGNITIVE_SERVICES_DOMAIN or hostname.endswith(
|
||||
f".{self.COGNITIVE_SERVICES_DOMAIN}"
|
||||
)
|
||||
|
||||
def _is_stt_endpoint(self, hostname: str) -> bool:
|
||||
return hostname == self.STT_SPEECH_DOMAIN or hostname.endswith(
|
||||
f".{self.STT_SPEECH_DOMAIN}"
|
||||
)
|
||||
|
||||
def _is_azure_openai_endpoint(self, hostname: str) -> bool:
|
||||
return hostname.endswith(".openai.azure.com")
|
||||
|
||||
def _extract_region_from_hostname(self, hostname: str, domain: str) -> str:
|
||||
if hostname.endswith(f".{domain}"):
|
||||
return hostname[: -len(f".{domain}")]
|
||||
return ""
|
||||
|
||||
def _build_stt_base_url(self, region: str) -> str:
|
||||
if region:
|
||||
return f"https://{region}.{self.STT_SPEECH_DOMAIN}"
|
||||
return f"https://{self.STT_SPEECH_DOMAIN}"
|
||||
|
||||
def _get_azure_response_format(self, response_format: Optional[str]) -> str:
|
||||
if response_format == "verbose_json":
|
||||
return "detailed"
|
||||
return "simple"
|
||||
|
||||
def _extract_text(self, response_json: Dict[str, Any]) -> str:
|
||||
if isinstance(response_json.get("DisplayText"), str):
|
||||
return response_json["DisplayText"]
|
||||
|
||||
nbest = response_json.get("NBest")
|
||||
if isinstance(nbest, list) and nbest:
|
||||
best = nbest[0]
|
||||
if isinstance(best, dict):
|
||||
return best.get("Display") or best.get("Lexical") or ""
|
||||
|
||||
return ""
|
||||
|
|
@ -239,7 +239,9 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
)
|
||||
|
||||
data = {"model": None, "messages": messages, **optional_params}
|
||||
elif litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(model=model):
|
||||
elif litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(
|
||||
model=litellm_params.get("base_model") or model
|
||||
):
|
||||
data = litellm.AzureOpenAIGPT5Config().transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
|
|
|
|||
|
|
@ -4,10 +4,10 @@ Support for o1 and o3 model families
|
|||
https://platform.openai.com/docs/guides/reasoning
|
||||
|
||||
Translations handled by LiteLLM:
|
||||
- modalities: image => drop param (if user opts in to dropping param)
|
||||
- role: system ==> translate to role 'user'
|
||||
- streaming => faked by LiteLLM
|
||||
- Tools, response_format => drop param (if user opts in to dropping param)
|
||||
- modalities: image => drop param (if user opts in to dropping param)
|
||||
- role: system ==> translate to role 'user'
|
||||
- streaming => faked by LiteLLM
|
||||
- Tools, response_format => drop param (if user opts in to dropping param)
|
||||
- Logprobs => drop param (if user opts in to dropping param)
|
||||
- Temperature => drop param (if user opts in to dropping param)
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,9 +1,16 @@
|
|||
from typing import Optional
|
||||
from urllib.parse import parse_qs, urlparse, urlunparse
|
||||
|
||||
from litellm.llms.azure.common_utils import BaseAzureLLM
|
||||
from litellm.llms.openai.containers.transformation import OpenAIContainerConfig
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
# Endpoint-specific path suffixes that may appear in a deployment's api_base
|
||||
# (e.g. the responses endpoint URL is stored as api_base for Azure models).
|
||||
# Strip these before building the containers URL so we always start from the
|
||||
# resource root (https://resource.cognitiveservices.azure.com).
|
||||
_AZURE_ENDPOINT_PATHS = ("/openai/responses",)
|
||||
|
||||
|
||||
class AzureContainerConfig(OpenAIContainerConfig):
|
||||
"""
|
||||
|
|
@ -27,6 +34,27 @@ class AzureContainerConfig(OpenAIContainerConfig):
|
|||
litellm_params=GenericLiteLLMParams(api_key=api_key),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _normalize_api_base(api_base: Optional[str]) -> Optional[str]:
|
||||
"""Strip endpoint-specific path suffixes from api_base to get the resource root."""
|
||||
if not api_base:
|
||||
return api_base
|
||||
parsed = urlparse(api_base)
|
||||
path = parsed.path.rstrip("/")
|
||||
for ep in _AZURE_ENDPOINT_PATHS:
|
||||
if path.endswith(ep):
|
||||
return urlunparse(
|
||||
(parsed.scheme, parsed.netloc, path[: -len(ep)], "", "", "")
|
||||
)
|
||||
return api_base
|
||||
|
||||
@staticmethod
|
||||
def _extract_api_version(api_base: Optional[str]) -> Optional[str]:
|
||||
"""Return the api-version query param from api_base if present."""
|
||||
if not api_base:
|
||||
return None
|
||||
return parse_qs(urlparse(api_base).query).get("api-version", [None])[0]
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
|
|
@ -39,10 +67,19 @@ class AzureContainerConfig(OpenAIContainerConfig):
|
|||
{endpoint}/openai/v1/containers
|
||||
when api_version is 'v1', 'latest', or 'preview'; otherwise:
|
||||
{endpoint}/openai/containers
|
||||
|
||||
The deployment's api_base may be the responses endpoint URL
|
||||
(e.g. .../openai/responses?api-version=2025-04-01-preview). We
|
||||
prefer the api-version embedded there over the deployment's
|
||||
api_version field, which may point to an older chat API version.
|
||||
"""
|
||||
effective_params = dict(litellm_params)
|
||||
api_version_from_base = self._extract_api_version(api_base)
|
||||
if api_version_from_base:
|
||||
effective_params["api_version"] = api_version_from_base
|
||||
return BaseAzureLLM._get_base_azure_url(
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
api_base=self._normalize_api_base(api_base),
|
||||
litellm_params=effective_params,
|
||||
route="/openai/containers",
|
||||
default_api_version="v1",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Transformation logic from OpenAI /v1/embeddings format to Azure AI Cohere's /v1/embed.
|
||||
Transformation logic from OpenAI /v1/embeddings format to Azure AI Cohere's /v1/embed.
|
||||
|
||||
Why separate file? Make it easy to see how transformation works
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Translate between Cohere's `/rerank` format and Azure AI's `/rerank` format.
|
||||
Translate between Cohere's `/rerank` format and Azure AI's `/rerank` format.
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
|
|
|
|||
0
litellm/llms/base_llm/agents/__init__.py
Normal file
0
litellm/llms/base_llm/agents/__init__.py
Normal file
165
litellm/llms/base_llm/agents/transformation.py
Normal file
165
litellm/llms/base_llm/agents/transformation.py
Normal file
|
|
@ -0,0 +1,165 @@
|
|||
"""
|
||||
Base transformation class for provider-side Agents API.
|
||||
|
||||
Providers that have a native agents CRUD API (e.g. Gemini v1beta/agents)
|
||||
subclass BaseAgentsAPIConfig and implement the abstract methods.
|
||||
|
||||
The HTTP calls are handled by AgentsHTTPHandler — this class is pure
|
||||
transform logic (same separation as BaseInteractionsAPIConfig /
|
||||
InteractionsHTTPHandler).
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Dict, Optional, Tuple, Union
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.types.agents import (
|
||||
AgentCreateResponse,
|
||||
AgentDeleteResult,
|
||||
AgentListResponse,
|
||||
AgentVersionsResponse,
|
||||
)
|
||||
|
||||
|
||||
class BaseAgentsAPIConfig(ABC):
|
||||
"""
|
||||
Minimal interface for providers that expose a native agents CRUD API.
|
||||
"""
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# CREATE #
|
||||
# ------------------------------------------------------------------ #
|
||||
|
||||
@abstractmethod
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
litellm_params: Dict[str, Any],
|
||||
) -> str:
|
||||
"""Return the full URL for POST /agents (create)."""
|
||||
|
||||
@abstractmethod
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Dict[str, str],
|
||||
litellm_params: Dict[str, Any],
|
||||
) -> Dict[str, str]:
|
||||
"""Validate credentials and return auth headers."""
|
||||
|
||||
@abstractmethod
|
||||
def transform_create_request(
|
||||
self,
|
||||
name: str,
|
||||
litellm_params: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
"""Map name + litellm_params to the provider's create-agent body."""
|
||||
|
||||
@abstractmethod
|
||||
def transform_create_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
name: str,
|
||||
) -> AgentCreateResponse:
|
||||
"""Parse create response. Raise on non-2xx."""
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# LIST #
|
||||
# ------------------------------------------------------------------ #
|
||||
|
||||
@abstractmethod
|
||||
def transform_list_request(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
litellm_params: Dict[str, Any],
|
||||
) -> Tuple[str, Dict[str, Any]]:
|
||||
"""Return (url, query_params) for GET /agents."""
|
||||
|
||||
@abstractmethod
|
||||
def transform_list_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
) -> AgentListResponse:
|
||||
"""Parse list-agents response. Raise on non-2xx."""
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# GET #
|
||||
# ------------------------------------------------------------------ #
|
||||
|
||||
@abstractmethod
|
||||
def transform_get_request(
|
||||
self,
|
||||
name: str,
|
||||
api_base: Optional[str],
|
||||
litellm_params: Dict[str, Any],
|
||||
) -> Tuple[str, Dict[str, Any]]:
|
||||
"""Return (url, query_params) for GET /agents/{name}."""
|
||||
|
||||
@abstractmethod
|
||||
def transform_get_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
name: str,
|
||||
) -> AgentCreateResponse:
|
||||
"""Parse get-agent response. Raise on non-2xx."""
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# DELETE #
|
||||
# ------------------------------------------------------------------ #
|
||||
|
||||
@abstractmethod
|
||||
def transform_delete_request(
|
||||
self,
|
||||
name: str,
|
||||
api_base: Optional[str],
|
||||
litellm_params: Dict[str, Any],
|
||||
) -> str:
|
||||
"""Return the URL for DELETE /agents/{name}."""
|
||||
|
||||
@abstractmethod
|
||||
def transform_delete_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
name: str,
|
||||
) -> AgentDeleteResult:
|
||||
"""Parse delete-agent response. Raise on non-2xx."""
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# LIST VERSIONS #
|
||||
# ------------------------------------------------------------------ #
|
||||
|
||||
@abstractmethod
|
||||
def transform_list_versions_request(
|
||||
self,
|
||||
name: str,
|
||||
api_base: Optional[str],
|
||||
litellm_params: Dict[str, Any],
|
||||
) -> Tuple[str, Dict[str, Any]]:
|
||||
"""Return (url, query_params) for GET /agents/{name}/versions."""
|
||||
|
||||
@abstractmethod
|
||||
def transform_list_versions_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
name: str,
|
||||
) -> AgentVersionsResponse:
|
||||
"""Parse list-versions response. Raise on non-2xx."""
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# ERROR HANDLING #
|
||||
# ------------------------------------------------------------------ #
|
||||
|
||||
def get_error_class(
|
||||
self,
|
||||
error_message: str,
|
||||
status_code: int,
|
||||
headers: Union[dict, httpx.Headers],
|
||||
) -> Exception:
|
||||
"""Map HTTP error status codes to provider-specific exceptions."""
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
|
||||
return BaseLLMException(
|
||||
status_code=status_code,
|
||||
message=error_message,
|
||||
headers=headers,
|
||||
)
|
||||
|
|
@ -54,6 +54,7 @@ class OCRUsageInfo(LiteLLMPydanticObjectBase):
|
|||
"""Usage information from OCR response."""
|
||||
|
||||
pages_processed: Optional[int] = None
|
||||
credits: Optional[float] = None
|
||||
doc_size_bytes: Optional[int] = None
|
||||
|
||||
model_config = {"extra": "allow"}
|
||||
|
|
|
|||
|
|
@ -44,6 +44,12 @@ else:
|
|||
# (e.g. "us-east-1", "eu-west-2", "us-gov-west-1", "cn-north-1").
|
||||
_VALID_AWS_REGION_PATTERN = re.compile(r"\A[a-z0-9-]+\Z")
|
||||
|
||||
# Regional STS hostnames, e.g. sts.eu-west-1.amazonaws.com or
|
||||
# vpce-xxx.sts.eu-west-1.vpce.amazonaws.com
|
||||
_STS_REGION_FROM_ENDPOINT_PATTERN = re.compile(
|
||||
r"(?:^|\.)sts(?:-fips)?\.([a-z0-9-]+)\.(?:amazonaws\.com(?:\.cn)?|vpce\.amazonaws\.com)"
|
||||
)
|
||||
|
||||
|
||||
class Boto3CredentialsInfo(BaseModel):
|
||||
credentials: Credentials
|
||||
|
|
@ -450,6 +456,24 @@ class BaseAWSLLM:
|
|||
model_id = BaseAWSLLM.encode_model_id(model_id=model_id)
|
||||
else:
|
||||
model_id = model
|
||||
# Strip LiteLLM routing prefixes (e.g. "bedrock/", "invoke/",
|
||||
# "bedrock/invoke/", "bedrock/converse/") that are not part of the
|
||||
# actual Bedrock model ID. The converse path already does this; the
|
||||
# invoke path must do the same so that ARN models such as
|
||||
# bedrock/arn:aws:bedrock:…:inference-profile/global.anthropic.…
|
||||
# are not forwarded verbatim to the Bedrock API, which would produce
|
||||
# a malformed URL and cause botocore's EventStreamBuffer to receive
|
||||
# a JSON error body instead of a binary event-stream — surfaced as a
|
||||
# misleading ChecksumMismatch (0x223a7b22 == ':{"').
|
||||
# Use strip_bedrock_routing_prefix (no break) so compound prefixes
|
||||
# like "bedrock/invoke/arn:..." are fully stripped in one call.
|
||||
from litellm.llms.bedrock.common_utils import strip_bedrock_routing_prefix
|
||||
|
||||
model_id = strip_bedrock_routing_prefix(model_id)
|
||||
# URL-encode ARNs so colons and slashes are safe in the URL path.
|
||||
if model_id.startswith("arn:"):
|
||||
model_id = BaseAWSLLM.encode_model_id(model_id=model_id)
|
||||
return model_id
|
||||
|
||||
model_id = model_id.replace("invoke/", "", 1)
|
||||
if provider == "llama" and "llama/" in model_id:
|
||||
|
|
@ -633,6 +657,40 @@ class BaseAWSLLM:
|
|||
"Region names must contain only lowercase letters, digits, and hyphens."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _parse_sts_region_from_endpoint(
|
||||
aws_sts_endpoint: Optional[str],
|
||||
) -> Optional[str]:
|
||||
"""Extract region from sts.{region}.amazonaws.com or vpce-x.sts.{region}.vpce.amazonaws.com."""
|
||||
if not aws_sts_endpoint:
|
||||
return None
|
||||
host = urllib.parse.urlparse(aws_sts_endpoint).hostname or ""
|
||||
match = _STS_REGION_FROM_ENDPOINT_PATTERN.search(host)
|
||||
return match.group(1) if match else None
|
||||
|
||||
@staticmethod
|
||||
def _resolve_sts_region(aws_sts_endpoint: Optional[str] = None) -> Optional[str]:
|
||||
"""STS signing region: parsed from aws_sts_endpoint else AWS_REGION / AWS_DEFAULT_REGION."""
|
||||
return (
|
||||
BaseAWSLLM._parse_sts_region_from_endpoint(aws_sts_endpoint)
|
||||
or os.getenv("AWS_REGION")
|
||||
or os.getenv("AWS_DEFAULT_REGION")
|
||||
)
|
||||
|
||||
def _build_sts_client_kwargs(
|
||||
self,
|
||||
aws_sts_endpoint: Optional[str] = None,
|
||||
ssl_verify: Optional[Union[bool, str]] = None,
|
||||
) -> dict:
|
||||
"""STS client kwargs with aligned endpoint_url and region_name (SigV4)."""
|
||||
kwargs: dict = {"verify": self._get_ssl_verify(ssl_verify)}
|
||||
if aws_sts_endpoint is not None:
|
||||
kwargs["endpoint_url"] = aws_sts_endpoint
|
||||
sts_region = self._resolve_sts_region(aws_sts_endpoint)
|
||||
if sts_region is not None:
|
||||
kwargs["region_name"] = sts_region
|
||||
return kwargs
|
||||
|
||||
def get_aws_region_name_for_non_llm_api_calls(
|
||||
self,
|
||||
aws_region_name: Optional[str] = None,
|
||||
|
|
@ -787,11 +845,6 @@ class BaseAWSLLM:
|
|||
f"IN Web Identity Token: {aws_web_identity_token} | Role Name: {aws_role_name} | Session Name: {aws_session_name}"
|
||||
)
|
||||
|
||||
if aws_sts_endpoint is None:
|
||||
sts_endpoint = f"https://sts.{aws_region_name}.amazonaws.com"
|
||||
else:
|
||||
sts_endpoint = aws_sts_endpoint
|
||||
|
||||
oidc_token = get_secret(aws_web_identity_token)
|
||||
|
||||
if oidc_token is None:
|
||||
|
|
@ -800,13 +853,13 @@ class BaseAWSLLM:
|
|||
status_code=401,
|
||||
)
|
||||
|
||||
sts_client_kwargs = self._build_sts_client_kwargs(
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
with tracer.trace("boto3.client(sts)"):
|
||||
sts_client = boto3.client(
|
||||
"sts",
|
||||
region_name=aws_region_name,
|
||||
endpoint_url=sts_endpoint,
|
||||
verify=self._get_ssl_verify(ssl_verify),
|
||||
)
|
||||
sts_client = boto3.client("sts", **sts_client_kwargs)
|
||||
|
||||
# https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRoleWithWebIdentity.html
|
||||
# https://boto3.amazonaws.com/v1/documentation/api/latest/reference/services/sts/client/assume_role_with_web_identity.html
|
||||
|
|
@ -847,7 +900,6 @@ class BaseAWSLLM:
|
|||
irsa_role_arn: str,
|
||||
aws_role_name: str,
|
||||
aws_session_name: str,
|
||||
region: str,
|
||||
web_identity_token_file: str,
|
||||
aws_external_id: Optional[str] = None,
|
||||
aws_sts_endpoint: Optional[str] = None,
|
||||
|
|
@ -862,12 +914,10 @@ class BaseAWSLLM:
|
|||
with open(web_identity_token_file, "r") as f:
|
||||
web_identity_token = f.read().strip()
|
||||
|
||||
irsa_sts_kwargs: dict = {
|
||||
"region_name": region,
|
||||
"verify": self._get_ssl_verify(ssl_verify),
|
||||
}
|
||||
if aws_sts_endpoint is not None:
|
||||
irsa_sts_kwargs["endpoint_url"] = aws_sts_endpoint
|
||||
irsa_sts_kwargs = self._build_sts_client_kwargs(
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
# Create an STS client without credentials
|
||||
with tracer.trace("boto3.client(sts) for manual IRSA"):
|
||||
|
|
@ -924,7 +974,6 @@ class BaseAWSLLM:
|
|||
self,
|
||||
aws_role_name: str,
|
||||
aws_session_name: str,
|
||||
region: str,
|
||||
aws_external_id: Optional[str] = None,
|
||||
aws_sts_endpoint: Optional[str] = None,
|
||||
ssl_verify: Optional[Union[bool, str]] = None,
|
||||
|
|
@ -932,12 +981,10 @@ class BaseAWSLLM:
|
|||
"""Handle same-account role assumption for IRSA."""
|
||||
import boto3
|
||||
|
||||
irsa_sts_kwargs: dict = {
|
||||
"region_name": region,
|
||||
"verify": self._get_ssl_verify(ssl_verify),
|
||||
}
|
||||
if aws_sts_endpoint is not None:
|
||||
irsa_sts_kwargs["endpoint_url"] = aws_sts_endpoint
|
||||
irsa_sts_kwargs = self._build_sts_client_kwargs(
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
verbose_logger.debug("Same account role assumption, using automatic IRSA")
|
||||
with tracer.trace("boto3.client(sts) with automatic IRSA"):
|
||||
|
|
@ -1010,12 +1057,6 @@ class BaseAWSLLM:
|
|||
web_identity_token_file = os.getenv("AWS_WEB_IDENTITY_TOKEN_FILE")
|
||||
irsa_role_arn = os.getenv("AWS_ROLE_ARN")
|
||||
|
||||
region = (
|
||||
aws_region_name
|
||||
or os.getenv("AWS_REGION")
|
||||
or os.getenv("AWS_DEFAULT_REGION")
|
||||
)
|
||||
|
||||
# If we have IRSA environment variables and no explicit credentials,
|
||||
# we need to use the web identity token flow
|
||||
if (
|
||||
|
|
@ -1031,16 +1072,12 @@ class BaseAWSLLM:
|
|||
)
|
||||
|
||||
try:
|
||||
# Use passed-in region when set, else env, else default (align with AssumeRole path)
|
||||
region = region or "us-east-1"
|
||||
|
||||
# Check if we need to do cross-account role assumption
|
||||
if aws_role_name != irsa_role_arn:
|
||||
sts_response = self._handle_irsa_cross_account(
|
||||
irsa_role_arn,
|
||||
aws_role_name,
|
||||
aws_session_name,
|
||||
region,
|
||||
web_identity_token_file,
|
||||
aws_external_id,
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
|
|
@ -1050,7 +1087,6 @@ class BaseAWSLLM:
|
|||
sts_response = self._handle_irsa_same_account(
|
||||
aws_role_name,
|
||||
aws_session_name,
|
||||
region,
|
||||
aws_external_id,
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
|
|
@ -1074,11 +1110,10 @@ class BaseAWSLLM:
|
|||
|
||||
# In EKS/IRSA environments, use ambient credentials (no explicit keys needed)
|
||||
# This allows the web identity token to work automatically
|
||||
sts_client_kwargs: dict = {"verify": self._get_ssl_verify(ssl_verify)}
|
||||
if region is not None:
|
||||
sts_client_kwargs["region_name"] = region
|
||||
if aws_sts_endpoint is not None:
|
||||
sts_client_kwargs["endpoint_url"] = aws_sts_endpoint
|
||||
sts_client_kwargs = self._build_sts_client_kwargs(
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
if aws_access_key_id is None and aws_secret_access_key is None:
|
||||
with tracer.trace("boto3.client(sts)"):
|
||||
sts_client = boto3.client("sts", **sts_client_kwargs)
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ from typing import Any, Dict, List, Literal, Optional, Union, cast
|
|||
|
||||
from httpx import Headers, Response
|
||||
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
|
@ -263,9 +264,32 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
|
|||
cancelling_at=None,
|
||||
cancelled_at=None,
|
||||
request_counts=None,
|
||||
metadata=original_request.get("metadata", {}),
|
||||
metadata=self._get_openai_compatible_batch_metadata(
|
||||
original_request.get("metadata", {})
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_openai_compatible_batch_metadata(metadata: Any) -> Dict[str, str]:
|
||||
"""
|
||||
OpenAI Batch metadata only accepts string values.
|
||||
"""
|
||||
if not isinstance(metadata, dict):
|
||||
return {}
|
||||
|
||||
sanitized_metadata: Dict[str, str] = {}
|
||||
for key, value in metadata.items():
|
||||
if key == "standard_logging_guardrail_information" or value is None:
|
||||
continue
|
||||
|
||||
str_key = str(key)
|
||||
if isinstance(value, str):
|
||||
sanitized_metadata[str_key] = value
|
||||
else:
|
||||
sanitized_metadata[str_key] = safe_dumps(value)
|
||||
|
||||
return sanitized_metadata
|
||||
|
||||
def transform_retrieve_batch_request(
|
||||
self,
|
||||
batch_id: str,
|
||||
|
|
|
|||
|
|
@ -299,9 +299,9 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM):
|
|||
)
|
||||
|
||||
def _get_response_stream_shape(self):
|
||||
from litellm.llms.bedrock.common_utils import BEDROCK_RESPONSE_STREAM_SHAPE
|
||||
from litellm.llms.bedrock.common_utils import get_bedrock_response_stream_shape
|
||||
|
||||
return BEDROCK_RESPONSE_STREAM_SHAPE
|
||||
return get_bedrock_response_stream_shape()
|
||||
|
||||
def _extract_response_content(self, events: InvokeAgentEventList) -> str:
|
||||
"""Extract the final response content from parsed events."""
|
||||
|
|
|
|||
|
|
@ -68,9 +68,9 @@ from litellm.utils import CustomStreamWrapper, get_secret
|
|||
|
||||
from ..base_aws_llm import BaseAWSLLM
|
||||
from ..common_utils import (
|
||||
BEDROCK_RESPONSE_STREAM_SHAPE,
|
||||
BedrockError,
|
||||
ModelResponseIterator,
|
||||
get_bedrock_response_stream_shape,
|
||||
get_bedrock_tool_name,
|
||||
)
|
||||
|
||||
|
|
@ -1828,7 +1828,8 @@ class AWSEventStreamDecoder:
|
|||
yield self._chunk_parser(chunk_data=_data)
|
||||
|
||||
def _parse_message_from_event(self, event) -> Optional[str]:
|
||||
if BEDROCK_RESPONSE_STREAM_SHAPE is None:
|
||||
response_stream_shape = get_bedrock_response_stream_shape()
|
||||
if response_stream_shape is None:
|
||||
raise BedrockError(
|
||||
status_code=500,
|
||||
message=(
|
||||
|
|
@ -1837,9 +1838,7 @@ class AWSEventStreamDecoder:
|
|||
),
|
||||
)
|
||||
response_dict = event.to_response_dict()
|
||||
parsed_response = self.parser.parse(
|
||||
response_dict, BEDROCK_RESPONSE_STREAM_SHAPE
|
||||
)
|
||||
parsed_response = self.parser.parse(response_dict, response_stream_shape)
|
||||
|
||||
if response_dict["status_code"] != 200:
|
||||
decoded_body = response_dict["body"].decode()
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ from typing import TYPE_CHECKING, Any, List, Optional
|
|||
import httpx
|
||||
|
||||
from litellm.anthropic_beta_headers_manager import filter_and_transform_beta_headers
|
||||
from litellm.litellm_core_utils.litellm_logging import verbose_logger
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
convert_to_anthropic_image_obj,
|
||||
)
|
||||
|
|
@ -22,6 +23,7 @@ from litellm.llms.bedrock.common_utils import (
|
|||
from litellm.types.llms.anthropic import ANTHROPIC_TOOL_SEARCH_BETA_HEADER
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse
|
||||
from litellm.utils import _supports_factory
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
|
@ -169,6 +171,24 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
|
|||
anthropic_request.pop("model", None)
|
||||
anthropic_request.pop("stream", None)
|
||||
anthropic_request.pop("output_format", None)
|
||||
if not (
|
||||
_supports_factory(
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
key="supports_output_config",
|
||||
)
|
||||
or AnthropicConfig._model_supports_effort_param(model)
|
||||
):
|
||||
if anthropic_request.pop("output_config", None) is not None:
|
||||
verbose_logger.warning(
|
||||
"Bedrock Invoke: stripping unsupported `output_config` for "
|
||||
"model=%s — neither `supports_output_config` nor any "
|
||||
"`supports_*_reasoning_effort` flag is set in "
|
||||
"model_prices_and_context_window.json. Add the capability "
|
||||
"flag to the model JSON entry if this model accepts "
|
||||
"`output_config`.",
|
||||
model,
|
||||
)
|
||||
if "anthropic_version" not in anthropic_request:
|
||||
anthropic_request["anthropic_version"] = self.anthropic_version
|
||||
|
||||
|
|
|
|||
|
|
@ -4,7 +4,6 @@ import litellm
|
|||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
|
||||
CLAUDE_PLATFORM_SERVICE_NAME: Literal["aws-external-anthropic"] = (
|
||||
"aws-external-anthropic"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ from __future__ import annotations
|
|||
Common utilities used across bedrock chat/embedding/image generation
|
||||
"""
|
||||
|
||||
import functools
|
||||
import json
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union
|
||||
|
|
@ -963,10 +964,8 @@ def _load_bedrock_response_stream_shape():
|
|||
"""
|
||||
Load the ResponseStream shape from botocore's bundled bedrock-runtime schema.
|
||||
|
||||
Called once at module import time; the result is stored in
|
||||
``BEDROCK_RESPONSE_STREAM_SHAPE`` and reused for the process lifetime.
|
||||
Returns ``None`` if botocore is unavailable or the service model cannot be
|
||||
loaded, so the module still imports cleanly.
|
||||
loaded.
|
||||
"""
|
||||
try:
|
||||
from botocore.loaders import Loader
|
||||
|
|
@ -977,15 +976,22 @@ def _load_bedrock_response_stream_shape():
|
|||
return ServiceModel(service_dict).shape_for("ResponseStream")
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
"litellm: could not pre-load bedrock-runtime response stream shape "
|
||||
"litellm: could not load bedrock-runtime response stream shape "
|
||||
"— Bedrock event-stream decoding will be unavailable. Error: %s",
|
||||
e,
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
# Eagerly resolved once per process — avoids per-instance or per-request disk I/O.
|
||||
BEDROCK_RESPONSE_STREAM_SHAPE = _load_bedrock_response_stream_shape()
|
||||
@functools.lru_cache(maxsize=1)
|
||||
def get_bedrock_response_stream_shape():
|
||||
"""
|
||||
Lazily load and cache the bedrock-runtime ResponseStream shape for the process.
|
||||
|
||||
Avoids importing botocore (and logging warnings) unless Bedrock event-stream
|
||||
decoding is actually needed.
|
||||
"""
|
||||
return _load_bedrock_response_stream_shape()
|
||||
|
||||
|
||||
class BedrockEventStreamDecoderBase:
|
||||
|
|
@ -999,7 +1005,8 @@ class BedrockEventStreamDecoderBase:
|
|||
self.parser = EventStreamJSONParser()
|
||||
|
||||
def _parse_message_from_event(self, event) -> Optional[str]:
|
||||
if BEDROCK_RESPONSE_STREAM_SHAPE is None:
|
||||
response_stream_shape = get_bedrock_response_stream_shape()
|
||||
if response_stream_shape is None:
|
||||
raise BedrockError(
|
||||
status_code=500,
|
||||
message=(
|
||||
|
|
@ -1008,9 +1015,7 @@ class BedrockEventStreamDecoderBase:
|
|||
),
|
||||
)
|
||||
response_dict = event.to_response_dict()
|
||||
parsed_response = self.parser.parse(
|
||||
response_dict, BEDROCK_RESPONSE_STREAM_SHAPE
|
||||
)
|
||||
parsed_response = self.parser.parse(response_dict, response_stream_shape)
|
||||
|
||||
if response_dict["status_code"] != 200:
|
||||
decoded_body = response_dict["body"].decode()
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Transformation logic from OpenAI /v1/embeddings format to Bedrock Amazon Titan G1 /invoke format.
|
||||
Transformation logic from OpenAI /v1/embeddings format to Bedrock Amazon Titan G1 /invoke format.
|
||||
|
||||
Why separate file? Make it easy to see how transformation works
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Transformation logic from OpenAI /v1/embeddings format to Bedrock Cohere /invoke format.
|
||||
Transformation logic from OpenAI /v1/embeddings format to Bedrock Cohere /invoke format.
|
||||
|
||||
Why separate file? Make it easy to see how transformation works
|
||||
"""
|
||||
|
|
@ -22,7 +22,7 @@ class BedrockCohereEmbeddingConfig:
|
|||
) -> dict:
|
||||
for k, v in non_default_params.items():
|
||||
if k == "encoding_format":
|
||||
optional_params["embedding_types"] = v
|
||||
optional_params["embedding_types"] = v if isinstance(v, list) else [v]
|
||||
elif k == "dimensions":
|
||||
optional_params["output_dimension"] = v
|
||||
return optional_params
|
||||
|
|
|
|||
|
|
@ -45,6 +45,7 @@ from litellm.types.router import GenericLiteLLMParams
|
|||
from litellm.types.utils import GenericStreamingChunk
|
||||
from litellm.types.utils import GenericStreamingChunk as GChunk
|
||||
from litellm.types.utils import ModelResponseStream
|
||||
from litellm.utils import _supports_factory
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
|
@ -557,7 +558,29 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
anthropic_messages_request=anthropic_messages_request,
|
||||
)
|
||||
|
||||
# 5a. Remove `custom` field from tools (Bedrock doesn't support it)
|
||||
# 5a. Bedrock Invoke supports output_config (effort) for Claude 4.6+ models,
|
||||
# but older models do not — strip it to avoid request rejection.
|
||||
# Ref: https://github.com/BerriAI/litellm/issues/22797
|
||||
if not (
|
||||
_supports_factory(
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
key="supports_output_config",
|
||||
)
|
||||
or AnthropicConfig._model_supports_effort_param(model)
|
||||
):
|
||||
if anthropic_messages_request.pop("output_config", None) is not None:
|
||||
verbose_logger.warning(
|
||||
"Bedrock Invoke: stripping unsupported `output_config` for "
|
||||
"model=%s — neither `supports_output_config` nor any "
|
||||
"`supports_*_reasoning_effort` flag is set in "
|
||||
"model_prices_and_context_window.json. Add the capability "
|
||||
"flag to the model JSON entry if this model accepts "
|
||||
"`output_config`.",
|
||||
model,
|
||||
)
|
||||
|
||||
# 5b. Remove `custom` field from tools (Bedrock doesn't support it)
|
||||
# Claude Code sends `custom: {defer_loading: true}` on tool definitions,
|
||||
# which causes Bedrock to reject the request with "Extra inputs are not permitted"
|
||||
# Ref: https://github.com/BerriAI/litellm/issues/22847
|
||||
|
|
|
|||
|
|
@ -16,7 +16,6 @@ from litellm.secret_managers.main import get_secret_str
|
|||
|
||||
from ...openai_like.chat.transformation import OpenAILikeChatConfig
|
||||
|
||||
|
||||
BEDROCK_MANTLE_DEFAULT_REGION = "us-east-1"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,5 @@
|
|||
import json
|
||||
from typing import Any, Optional
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from litellm.constants import STREAM_SSE_DONE_STRING
|
||||
from litellm.exceptions import AuthenticationError
|
||||
from litellm.litellm_core_utils.core_helpers import process_response_headers
|
||||
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
|
||||
|
|
@ -9,13 +7,17 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
|
|||
)
|
||||
from litellm.llms.openai.common_utils import OpenAIError
|
||||
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
|
||||
from litellm.responses.sse_output_recovery import (
|
||||
parse_sse_json_chunk,
|
||||
record_output_item_chunk,
|
||||
record_output_text_chunk,
|
||||
)
|
||||
from litellm.types.llms.openai import (
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.utils import CustomStreamWrapper
|
||||
|
||||
from ..authenticator import Authenticator
|
||||
from ..common_utils import (
|
||||
|
|
@ -111,86 +113,139 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
raw_response: Any,
|
||||
logging_obj: Any,
|
||||
):
|
||||
content_type = (raw_response.headers or {}).get("content-type", "")
|
||||
body_text = raw_response.text or ""
|
||||
if "text/event-stream" not in content_type.lower():
|
||||
trimmed_body = body_text.lstrip()
|
||||
if not (
|
||||
trimmed_body.startswith("event:")
|
||||
or trimmed_body.startswith("data:")
|
||||
or "\nevent:" in body_text
|
||||
or "\ndata:" in body_text
|
||||
):
|
||||
return super().transform_response_api_response(
|
||||
model=model,
|
||||
raw_response=raw_response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
if not self._should_parse_as_sse(
|
||||
raw_response=raw_response, body_text=body_text
|
||||
):
|
||||
return super().transform_response_api_response(
|
||||
model=model,
|
||||
raw_response=raw_response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
logging_obj.post_call(
|
||||
original_response=raw_response.text,
|
||||
additional_args={"complete_input_dict": {}},
|
||||
)
|
||||
|
||||
completed_response = None
|
||||
error_message = None
|
||||
for chunk in body_text.splitlines():
|
||||
stripped_chunk = CustomStreamWrapper._strip_sse_data_from_chunk(chunk)
|
||||
if not stripped_chunk:
|
||||
continue
|
||||
stripped_chunk = stripped_chunk.strip()
|
||||
if not stripped_chunk:
|
||||
continue
|
||||
if stripped_chunk == STREAM_SSE_DONE_STRING:
|
||||
break
|
||||
try:
|
||||
parsed_chunk = json.loads(stripped_chunk)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
if not isinstance(parsed_chunk, dict):
|
||||
continue
|
||||
event_type = parsed_chunk.get("type")
|
||||
if event_type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED:
|
||||
response_payload = parsed_chunk.get("response")
|
||||
if isinstance(response_payload, dict):
|
||||
response_payload = dict(response_payload)
|
||||
if "created_at" in response_payload:
|
||||
response_payload["created_at"] = _safe_convert_created_field(
|
||||
response_payload["created_at"]
|
||||
)
|
||||
try:
|
||||
completed_response = ResponsesAPIResponse(**response_payload)
|
||||
except Exception:
|
||||
completed_response = ResponsesAPIResponse.model_construct(
|
||||
**response_payload
|
||||
)
|
||||
break
|
||||
if event_type in (
|
||||
ResponsesAPIStreamEvents.RESPONSE_FAILED,
|
||||
ResponsesAPIStreamEvents.ERROR,
|
||||
):
|
||||
error_obj = parsed_chunk.get("error") or (
|
||||
parsed_chunk.get("response") or {}
|
||||
).get("error")
|
||||
if error_obj is not None:
|
||||
if isinstance(error_obj, dict):
|
||||
error_message = error_obj.get("message") or str(error_obj)
|
||||
else:
|
||||
error_message = str(error_obj)
|
||||
|
||||
completed_response, error_message = self._extract_completed_response_from_sse(
|
||||
body_text=body_text
|
||||
)
|
||||
if completed_response is None:
|
||||
raise OpenAIError(
|
||||
message=error_message or raw_response.text,
|
||||
status_code=raw_response.status_code,
|
||||
)
|
||||
|
||||
self._attach_response_headers(
|
||||
completed_response=completed_response, raw_response=raw_response
|
||||
)
|
||||
return completed_response
|
||||
|
||||
def _should_parse_as_sse(self, raw_response: Any, body_text: str) -> bool:
|
||||
content_type = (raw_response.headers or {}).get("content-type", "")
|
||||
if "text/event-stream" in content_type.lower():
|
||||
return True
|
||||
trimmed_body = body_text.lstrip()
|
||||
return bool(
|
||||
trimmed_body.startswith("event:")
|
||||
or trimmed_body.startswith("data:")
|
||||
or "\nevent:" in body_text
|
||||
or "\ndata:" in body_text
|
||||
)
|
||||
|
||||
def _extract_completed_response_from_sse(
|
||||
self, body_text: str
|
||||
) -> tuple[Optional[ResponsesAPIResponse], Optional[str]]:
|
||||
completed_response = None
|
||||
error_message = None
|
||||
streamed_output_items: Dict[int, dict] = {}
|
||||
text_only_output_items: Dict[int, dict] = {}
|
||||
for chunk in body_text.splitlines():
|
||||
parsed_chunk = parse_sse_json_chunk(chunk)
|
||||
if parsed_chunk is None:
|
||||
continue
|
||||
|
||||
event_type = parsed_chunk.get("type")
|
||||
if event_type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE:
|
||||
record_output_item_chunk(
|
||||
parsed_chunk=parsed_chunk,
|
||||
output_items=streamed_output_items,
|
||||
)
|
||||
continue
|
||||
|
||||
if event_type == ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE:
|
||||
record_output_text_chunk(
|
||||
parsed_chunk=parsed_chunk,
|
||||
output_items=streamed_output_items,
|
||||
text_only_items=text_only_output_items,
|
||||
)
|
||||
continue
|
||||
|
||||
if event_type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED:
|
||||
# Real OUTPUT_ITEM_DONE events take precedence at any given
|
||||
# output_index, but text-only items at indices without a
|
||||
# matching OUTPUT_ITEM_DONE must still be preserved (e.g.
|
||||
# providers that emit only OUTPUT_TEXT_DONE for some indices).
|
||||
merged_items: Dict[int, dict] = {**text_only_output_items}
|
||||
merged_items.update(streamed_output_items)
|
||||
completed_response = self._build_completed_response_from_chunk(
|
||||
parsed_chunk=parsed_chunk,
|
||||
streamed_output_items=merged_items,
|
||||
)
|
||||
break
|
||||
|
||||
if event_type in (
|
||||
ResponsesAPIStreamEvents.RESPONSE_FAILED,
|
||||
ResponsesAPIStreamEvents.ERROR,
|
||||
):
|
||||
extracted_error = self._extract_error_message(parsed_chunk)
|
||||
if extracted_error is not None:
|
||||
error_message = extracted_error
|
||||
|
||||
return completed_response, error_message
|
||||
|
||||
def _build_completed_response_from_chunk(
|
||||
self, parsed_chunk: Dict[str, Any], streamed_output_items: Dict[int, dict]
|
||||
) -> Optional[ResponsesAPIResponse]:
|
||||
response_payload = parsed_chunk.get("response")
|
||||
if not isinstance(response_payload, dict):
|
||||
return None
|
||||
response_payload = dict(response_payload)
|
||||
if not response_payload.get("output") and streamed_output_items:
|
||||
response_payload["output"] = [
|
||||
item for _, item in sorted(streamed_output_items.items())
|
||||
]
|
||||
if "created_at" in response_payload:
|
||||
response_payload["created_at"] = _safe_convert_created_field(
|
||||
response_payload["created_at"]
|
||||
)
|
||||
try:
|
||||
return ResponsesAPIResponse(**response_payload)
|
||||
except Exception:
|
||||
return ResponsesAPIResponse.model_construct(**response_payload)
|
||||
|
||||
def _extract_error_message(self, parsed_chunk: Dict[str, Any]) -> Optional[str]:
|
||||
error_obj = parsed_chunk.get("error") or (
|
||||
parsed_chunk.get("response") or {}
|
||||
).get("error")
|
||||
if error_obj is None:
|
||||
return None
|
||||
if isinstance(error_obj, dict):
|
||||
return error_obj.get("message") or str(error_obj)
|
||||
return str(error_obj)
|
||||
|
||||
def _attach_response_headers(
|
||||
self,
|
||||
completed_response: ResponsesAPIResponse,
|
||||
raw_response: Any,
|
||||
) -> None:
|
||||
raw_headers = dict(raw_response.headers)
|
||||
processed_headers = process_response_headers(raw_headers)
|
||||
if not hasattr(completed_response, "_hidden_params"):
|
||||
setattr(completed_response, "_hidden_params", {})
|
||||
completed_response._hidden_params["additional_headers"] = processed_headers
|
||||
completed_response._hidden_params["headers"] = raw_headers
|
||||
return completed_response
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Legacy /v1/embedding handler for Bedrock Cohere.
|
||||
Legacy /v1/embedding handler for Bedrock Cohere.
|
||||
"""
|
||||
|
||||
import json
|
||||
|
|
|
|||
|
|
@ -110,15 +110,35 @@ class CohereEmbeddingConfig:
|
|||
additional_args={"complete_input_dict": data},
|
||||
original_response=response_json,
|
||||
)
|
||||
return self._populate_embedding_response(
|
||||
response_json=response_json,
|
||||
model_response=model_response,
|
||||
model=model,
|
||||
encoding=encoding,
|
||||
input=input,
|
||||
)
|
||||
|
||||
def _populate_embedding_response(
|
||||
self,
|
||||
response_json: dict,
|
||||
model_response: EmbeddingResponse,
|
||||
model: str,
|
||||
encoding: Any,
|
||||
input: list,
|
||||
) -> EmbeddingResponse:
|
||||
"""
|
||||
response
|
||||
Parse a Cohere embed response body into an OpenAI-style EmbeddingResponse.
|
||||
|
||||
Split out from `_transform_response` so callers that already log
|
||||
`post_call` themselves (e.g. SageMaker's embedding handler) can reuse
|
||||
the parsing without triggering a second `post_call`.
|
||||
|
||||
Response shape:
|
||||
{
|
||||
'object': "list",
|
||||
'data': [
|
||||
|
||||
]
|
||||
'model',
|
||||
'usage'
|
||||
'data': [...],
|
||||
'model',
|
||||
'usage',
|
||||
}
|
||||
"""
|
||||
embeddings = response_json["embeddings"]
|
||||
|
|
@ -149,9 +169,6 @@ class CohereEmbeddingConfig:
|
|||
model_response.object = "list"
|
||||
model_response.data = output_data
|
||||
model_response.model = model
|
||||
input_tokens = 0
|
||||
for text in input:
|
||||
input_tokens += len(encoding.encode(text))
|
||||
|
||||
setattr(
|
||||
model_response,
|
||||
|
|
|
|||
|
|
@ -257,14 +257,19 @@ class GenericContainerHandler:
|
|||
returns_binary = endpoint_config.get("returns_binary", False)
|
||||
is_multipart = endpoint_config.get("is_multipart", False)
|
||||
|
||||
# An empty dict passed as `params` to httpx strips any existing query
|
||||
# string from the URL (e.g. ?api-version=...). Use None instead so
|
||||
# httpx leaves the URL's own query string intact.
|
||||
effective_params = query_params or None
|
||||
|
||||
try:
|
||||
if method == "GET":
|
||||
response = http_client.get(
|
||||
url=url, headers=headers, params=query_params
|
||||
url=url, headers=headers, params=effective_params
|
||||
)
|
||||
elif method == "DELETE":
|
||||
response = http_client.delete(
|
||||
url=url, headers=headers, params=query_params
|
||||
url=url, headers=headers, params=effective_params
|
||||
)
|
||||
elif method == "POST":
|
||||
if is_multipart and "file" in kwargs:
|
||||
|
|
@ -272,11 +277,11 @@ class GenericContainerHandler:
|
|||
kwargs["file"], headers
|
||||
)
|
||||
response = http_client.post(
|
||||
url=url, headers=headers, params=query_params, files=files
|
||||
url=url, headers=headers, params=effective_params, files=files
|
||||
)
|
||||
else:
|
||||
response = http_client.post(
|
||||
url=url, headers=headers, params=query_params
|
||||
url=url, headers=headers, params=effective_params
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unsupported HTTP method: {method}")
|
||||
|
|
@ -376,14 +381,19 @@ class GenericContainerHandler:
|
|||
returns_binary = endpoint_config.get("returns_binary", False)
|
||||
is_multipart = endpoint_config.get("is_multipart", False)
|
||||
|
||||
# An empty dict passed as `params` to httpx strips any existing query
|
||||
# string from the URL (e.g. ?api-version=...). Use None instead so
|
||||
# httpx leaves the URL's own query string intact.
|
||||
effective_params = query_params or None
|
||||
|
||||
try:
|
||||
if method == "GET":
|
||||
response = await http_client.get(
|
||||
url=url, headers=headers, params=query_params
|
||||
url=url, headers=headers, params=effective_params
|
||||
)
|
||||
elif method == "DELETE":
|
||||
response = await http_client.delete(
|
||||
url=url, headers=headers, params=query_params
|
||||
url=url, headers=headers, params=effective_params
|
||||
)
|
||||
elif method == "POST":
|
||||
if is_multipart and "file" in kwargs:
|
||||
|
|
@ -391,11 +401,11 @@ class GenericContainerHandler:
|
|||
kwargs["file"], headers
|
||||
)
|
||||
response = await http_client.post(
|
||||
url=url, headers=headers, params=query_params, files=files
|
||||
url=url, headers=headers, params=effective_params, files=files
|
||||
)
|
||||
else:
|
||||
response = await http_client.post(
|
||||
url=url, headers=headers, params=query_params
|
||||
url=url, headers=headers, params=effective_params
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unsupported HTTP method: {method}")
|
||||
|
|
|
|||
|
|
@ -890,6 +890,18 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
)
|
||||
|
||||
# Some providers (e.g. OCI) require request signing after the body is built.
|
||||
# The default BaseConfig.sign_request returns (headers, None) — a no-op for
|
||||
# providers that don't need signing.
|
||||
headers, signed_body = provider_config.sign_request(
|
||||
headers=headers,
|
||||
optional_params=optional_params,
|
||||
request_data=data,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=input,
|
||||
|
|
@ -916,6 +928,7 @@ class BaseLLMHTTPHandler:
|
|||
client=client,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
signed_body=signed_body,
|
||||
)
|
||||
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
|
|
@ -926,12 +939,20 @@ class BaseLLMHTTPHandler:
|
|||
sync_httpx_client = client
|
||||
|
||||
try:
|
||||
response = sync_httpx_client.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
data=json.dumps(data),
|
||||
timeout=timeout,
|
||||
)
|
||||
if signed_body is not None:
|
||||
response = sync_httpx_client.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
data=signed_body,
|
||||
timeout=timeout,
|
||||
)
|
||||
else:
|
||||
response = sync_httpx_client.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
data=json.dumps(data),
|
||||
timeout=timeout,
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
|
|
@ -964,6 +985,7 @@ class BaseLLMHTTPHandler:
|
|||
api_key: Optional[str] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
|
||||
signed_body: Optional[bytes] = None,
|
||||
) -> EmbeddingResponse:
|
||||
if client is None or not isinstance(client, AsyncHTTPHandler):
|
||||
async_httpx_client = get_async_httpx_client(
|
||||
|
|
@ -974,12 +996,20 @@ class BaseLLMHTTPHandler:
|
|||
async_httpx_client = client
|
||||
|
||||
try:
|
||||
response = await async_httpx_client.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
json=request_data,
|
||||
timeout=timeout,
|
||||
)
|
||||
if signed_body is not None:
|
||||
response = await async_httpx_client.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
data=signed_body,
|
||||
timeout=timeout,
|
||||
)
|
||||
else:
|
||||
response = await async_httpx_client.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
json=request_data,
|
||||
timeout=timeout,
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=provider_config)
|
||||
|
||||
|
|
@ -1177,6 +1207,8 @@ class BaseLLMHTTPHandler:
|
|||
|
||||
data = transformed_result.data
|
||||
files = transformed_result.files
|
||||
if transformed_result.content_type is not None:
|
||||
headers["Content-Type"] = transformed_result.content_type
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
@ -1409,6 +1441,8 @@ class BaseLLMHTTPHandler:
|
|||
document=document,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
# All providers return OCRRequestData
|
||||
|
|
@ -1477,6 +1511,8 @@ class BaseLLMHTTPHandler:
|
|||
document=document,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
# All providers return OCRRequestData
|
||||
|
|
@ -1852,7 +1888,9 @@ class BaseLLMHTTPHandler:
|
|||
async_httpx_client: AsyncHTTPHandler,
|
||||
request_url: str,
|
||||
headers: dict,
|
||||
signed_json_body: Optional[bytes],
|
||||
# str when the caller passes a pre-serialized (unsigned) body to avoid
|
||||
# re-dumping; bytes when a provider signed the request (e.g. Bedrock).
|
||||
signed_json_body: Optional[Union[str, bytes]],
|
||||
request_body: dict,
|
||||
stream: bool,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
|
|
@ -2043,8 +2081,18 @@ class BaseLLMHTTPHandler:
|
|||
model=model,
|
||||
)
|
||||
|
||||
# The request body was serialized once for the pre-call log input and
|
||||
# again for the wire (json.dumps is O(payload), large for long-context
|
||||
# Claude Code history). Serialize once and reuse for both. Only when
|
||||
# the provider didn't sign the request (sign_request no-op for the
|
||||
# native anthropic path -> signed_json_body is None); signed providers
|
||||
# (e.g. Bedrock) keep their signed body untouched. The HTTP-error
|
||||
# retry path mutates + re-signs the body, so it still re-serializes
|
||||
# internally -- this only deduplicates the success path.
|
||||
request_body_json = json.dumps(request_body)
|
||||
|
||||
logging_obj.pre_call(
|
||||
input=[{"role": "user", "content": json.dumps(request_body)}],
|
||||
input=[{"role": "user", "content": request_body_json}],
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": request_body,
|
||||
|
|
@ -2057,7 +2105,9 @@ class BaseLLMHTTPHandler:
|
|||
async_httpx_client=async_httpx_client,
|
||||
request_url=request_url,
|
||||
headers=headers,
|
||||
signed_json_body=signed_json_body,
|
||||
signed_json_body=(
|
||||
signed_json_body if signed_json_body is not None else request_body_json
|
||||
),
|
||||
request_body=request_body,
|
||||
stream=stream or False,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -2079,6 +2129,14 @@ class BaseLLMHTTPHandler:
|
|||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
if not self._has_agentic_completion_hook(logging_obj):
|
||||
# No callback overrides async_should_run_agentic_loop, so the
|
||||
# agentic wrapper's only effect would be buffering every chunk
|
||||
# and rebuilding the response from SSE at end-of-stream to call
|
||||
# hooks that all return (False, {}). Stream through directly and
|
||||
# skip that per-chunk + end-of-stream overhead.
|
||||
return completion_stream
|
||||
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
|
||||
AgenticAnthropicStreamingIterator,
|
||||
)
|
||||
|
|
@ -4586,6 +4644,51 @@ class BaseLLMHTTPHandler:
|
|||
fingerprints = list(kwargs.get("_agentic_loop_fingerprints", []) or [])
|
||||
return depth, max(max_loops, 1), fingerprints
|
||||
|
||||
@staticmethod
|
||||
def _has_agentic_completion_hook(logging_obj: Any) -> bool:
|
||||
"""
|
||||
True if any registered callback actually overrides
|
||||
``async_should_run_agentic_loop`` (the gate every agentic hook goes
|
||||
through). The base ``CustomLogger`` implementation returns
|
||||
``(False, {})``, so when nothing overrides it the agentic
|
||||
post-processing is a guaranteed no-op and the streaming wrapper that
|
||||
buffers + rebuilds the whole response from SSE just to call it can be
|
||||
skipped entirely.
|
||||
|
||||
Function-identity comparison (not a leaf ``__dict__`` check) so an
|
||||
override inherited through any intermediate class is still detected --
|
||||
a false negative here would silently disable agentic features.
|
||||
|
||||
String entries in ``litellm.callbacks`` (e.g. ``"datadog"``) are
|
||||
resolved to their ``CustomLogger`` instance via
|
||||
``get_custom_logger_compatible_class`` -- same pattern as
|
||||
``ProxyLogging._callback_capabilities`` -- so a string-registered
|
||||
agentic callback is detected too.
|
||||
"""
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
get_custom_logger_compatible_class,
|
||||
)
|
||||
|
||||
base_func = CustomLogger.async_should_run_agentic_loop
|
||||
callbacks = litellm.callbacks + (
|
||||
getattr(logging_obj, "dynamic_success_callbacks", None) or []
|
||||
)
|
||||
for cb in callbacks:
|
||||
if isinstance(cb, str):
|
||||
resolved = get_custom_logger_compatible_class(cb) # type: ignore[arg-type]
|
||||
if resolved is None:
|
||||
continue
|
||||
cb = resolved
|
||||
if not isinstance(cb, CustomLogger):
|
||||
continue
|
||||
cb_func = getattr(type(cb), "async_should_run_agentic_loop", base_func)
|
||||
if getattr(cb_func, "__func__", cb_func) is not getattr(
|
||||
base_func, "__func__", base_func
|
||||
):
|
||||
return True
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _check_agentic_loop_safety(
|
||||
tool_calls: Any,
|
||||
|
|
@ -7834,7 +7937,7 @@ class BaseLLMHTTPHandler:
|
|||
response = sync_httpx_client.get(
|
||||
url=url,
|
||||
headers=headers,
|
||||
params=params,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_list_response(
|
||||
|
|
@ -7911,7 +8014,7 @@ class BaseLLMHTTPHandler:
|
|||
response = await async_httpx_client.get(
|
||||
url=url,
|
||||
headers=headers,
|
||||
params=params,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_list_response(
|
||||
|
|
@ -8001,7 +8104,7 @@ class BaseLLMHTTPHandler:
|
|||
response = sync_httpx_client.get(
|
||||
url=url,
|
||||
headers=headers,
|
||||
params=params,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_retrieve_response(
|
||||
|
|
@ -8078,7 +8181,7 @@ class BaseLLMHTTPHandler:
|
|||
response = await async_httpx_client.get(
|
||||
url=url,
|
||||
headers=headers,
|
||||
params=params,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_retrieve_response(
|
||||
|
|
@ -8168,7 +8271,7 @@ class BaseLLMHTTPHandler:
|
|||
response = sync_httpx_client.delete(
|
||||
url=url,
|
||||
headers=headers,
|
||||
params=params,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_delete_response(
|
||||
|
|
@ -8245,7 +8348,7 @@ class BaseLLMHTTPHandler:
|
|||
response = await async_httpx_client.delete(
|
||||
url=url,
|
||||
headers=headers,
|
||||
params=params,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_delete_response(
|
||||
|
|
@ -8341,7 +8444,7 @@ class BaseLLMHTTPHandler:
|
|||
response = sync_httpx_client.get(
|
||||
url=url,
|
||||
headers=headers,
|
||||
params=params,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_file_list_response(
|
||||
|
|
@ -8420,7 +8523,7 @@ class BaseLLMHTTPHandler:
|
|||
response = await async_httpx_client.get(
|
||||
url=url,
|
||||
headers=headers,
|
||||
params=params,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_file_list_response(
|
||||
|
|
@ -8508,7 +8611,7 @@ class BaseLLMHTTPHandler:
|
|||
response = sync_httpx_client.get(
|
||||
url=url,
|
||||
headers=headers,
|
||||
params=params,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_file_content_response(
|
||||
|
|
@ -8584,7 +8687,7 @@ class BaseLLMHTTPHandler:
|
|||
response = await async_httpx_client.get(
|
||||
url=url,
|
||||
headers=headers,
|
||||
params=params,
|
||||
params=params or None,
|
||||
)
|
||||
|
||||
return container_provider_config.transform_container_file_content_response(
|
||||
|
|
|
|||
|
|
@ -13,7 +13,6 @@ from typing import Tuple
|
|||
|
||||
import httpx
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pre-built response templates
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
28
litellm/llms/dashscope/common_utils.py
Normal file
28
litellm/llms/dashscope/common_utils.py
Normal file
|
|
@ -0,0 +1,28 @@
|
|||
"""
|
||||
Common utilities for the DashScope LLM provider.
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
|
||||
|
||||
class DashScopeError(BaseLLMException):
|
||||
"""Exception class for DashScope provider errors."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
status_code: int,
|
||||
message: str,
|
||||
headers: Optional[httpx.Headers] = None,
|
||||
):
|
||||
self.status_code = status_code
|
||||
self.message = message
|
||||
self.headers = headers or httpx.Headers()
|
||||
super().__init__(
|
||||
status_code=status_code,
|
||||
message=message,
|
||||
headers=dict(self.headers),
|
||||
)
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Cost calculator for Dashscope Chat models.
|
||||
Cost calculator for Dashscope Chat models.
|
||||
|
||||
Handles tiered pricing and prompt caching scenarios.
|
||||
"""
|
||||
|
|
|
|||
7
litellm/llms/dashscope/embed/__init__.py
Normal file
7
litellm/llms/dashscope/embed/__init__.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
"""
|
||||
DashScope Embedding Module
|
||||
"""
|
||||
|
||||
from .transformation import DashScopeEmbeddingConfig
|
||||
|
||||
__all__ = ["DashScopeEmbeddingConfig"]
|
||||
191
litellm/llms/dashscope/embed/transformation.py
Normal file
191
litellm/llms/dashscope/embed/transformation.py
Normal file
|
|
@ -0,0 +1,191 @@
|
|||
"""
|
||||
Transformation logic from OpenAI /v1/embeddings format to DashScope's /v1/embeddings format.
|
||||
|
||||
Supports
|
||||
- text-embedding-v4
|
||||
- text-embedding-v3
|
||||
|
||||
Endpoint
|
||||
- https://dashscope.aliyuncs.com/compatible-mode/v1/embeddings
|
||||
|
||||
Docs - https://help.aliyun.com/zh/model-studio/text-embedding-synchronous-api
|
||||
"""
|
||||
|
||||
from typing import List, Optional, Union
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues
|
||||
from litellm.types.utils import EmbeddingResponse, Usage
|
||||
|
||||
from ..common_utils import DashScopeError
|
||||
|
||||
DEFAULT_API_BASE = "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||
|
||||
|
||||
class DashScopeEmbeddingConfig(BaseEmbeddingConfig):
|
||||
"""
|
||||
Reference: https://help.aliyun.com/zh/model-studio/text-embedding-synchronous-api
|
||||
|
||||
DashScope exposes an OpenAI-compatible /v1/embeddings endpoint, so the
|
||||
request and response shapes are nearly identical to OpenAI's.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
# DashScope's compatible-mode embeddings API accepts the same params as OpenAI.
|
||||
# `dimensions` / `encoding_format` are only honored by text-embedding-v3 / v4;
|
||||
# earlier versions silently ignore them server-side.
|
||||
return ["dimensions", "encoding_format", "user"]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool = False,
|
||||
) -> dict:
|
||||
supported = self.get_supported_openai_params(model)
|
||||
for k, v in non_default_params.items():
|
||||
if v is None:
|
||||
continue
|
||||
if k in supported:
|
||||
optional_params[k] = v
|
||||
# unsupported params are dropped when drop_params=True;
|
||||
# the upstream _check_valid_arg already raised UnsupportedParamsError
|
||||
# for drop_params=False before this method is called.
|
||||
return optional_params
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
if api_key is None:
|
||||
api_key = get_secret_str("DASHSCOPE_API_KEY")
|
||||
if api_key is None:
|
||||
raise ValueError(
|
||||
"DashScope API key is required. Set 'DASHSCOPE_API_KEY' env var or pass api_key explicitly."
|
||||
)
|
||||
default_headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
}
|
||||
return {**default_headers, **headers}
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
base = api_base or get_secret_str("DASHSCOPE_API_BASE") or DEFAULT_API_BASE
|
||||
base = base.rstrip("/")
|
||||
if base.endswith("/embeddings"):
|
||||
return base
|
||||
return f"{base}/embeddings"
|
||||
|
||||
def transform_embedding_request(
|
||||
self,
|
||||
model: str,
|
||||
input: AllEmbeddingInputValues,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
data: dict = {
|
||||
"model": model,
|
||||
"input": input,
|
||||
}
|
||||
for key in ("dimensions", "encoding_format", "user"):
|
||||
value = optional_params.get(key)
|
||||
if value is not None:
|
||||
data[key] = value
|
||||
return data
|
||||
|
||||
def transform_embedding_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: EmbeddingResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
api_key: Optional[str],
|
||||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
) -> EmbeddingResponse:
|
||||
try:
|
||||
response_json = raw_response.json()
|
||||
except Exception as e:
|
||||
raise DashScopeError(
|
||||
status_code=raw_response.status_code,
|
||||
message=f"Failed to parse DashScope response as JSON: {str(e)}",
|
||||
)
|
||||
|
||||
logging_obj.post_call(
|
||||
input=request_data.get("input"),
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": request_data},
|
||||
original_response=response_json,
|
||||
)
|
||||
|
||||
if "error" in response_json:
|
||||
error = response_json["error"]
|
||||
message = (
|
||||
error.get("message", str(error))
|
||||
if isinstance(error, dict)
|
||||
else str(error)
|
||||
)
|
||||
raise DashScopeError(
|
||||
status_code=raw_response.status_code,
|
||||
message=message,
|
||||
)
|
||||
|
||||
model_response.object = "list"
|
||||
model_response.data = response_json.get("data", [])
|
||||
model_response.model = response_json.get("model", model)
|
||||
|
||||
usage = response_json.get("usage") or {}
|
||||
prompt_tokens = usage.get("prompt_tokens", 0)
|
||||
total_tokens = usage.get("total_tokens", prompt_tokens)
|
||||
setattr(
|
||||
model_response,
|
||||
"usage",
|
||||
Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=0,
|
||||
total_tokens=total_tokens,
|
||||
),
|
||||
)
|
||||
|
||||
if "id" in response_json:
|
||||
setattr(model_response, "id", response_json["id"])
|
||||
|
||||
return model_response
|
||||
|
||||
def get_error_class(
|
||||
self,
|
||||
error_message: str,
|
||||
status_code: int,
|
||||
headers: Union[dict, httpx.Headers],
|
||||
) -> BaseLLMException:
|
||||
if isinstance(headers, dict):
|
||||
headers = httpx.Headers(headers)
|
||||
return DashScopeError(
|
||||
status_code=status_code,
|
||||
message=error_message,
|
||||
headers=headers,
|
||||
)
|
||||
7
litellm/llms/dashscope/rerank/__init__.py
Normal file
7
litellm/llms/dashscope/rerank/__init__.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
"""
|
||||
DashScope Rerank Module
|
||||
"""
|
||||
|
||||
from .transformation import DashScopeRerankConfig
|
||||
|
||||
__all__ = ["DashScopeRerankConfig"]
|
||||
241
litellm/llms/dashscope/rerank/transformation.py
Normal file
241
litellm/llms/dashscope/rerank/transformation.py
Normal file
|
|
@ -0,0 +1,241 @@
|
|||
"""
|
||||
Transformation logic for DashScope's OpenAI-compatible /v1/reranks API.
|
||||
|
||||
Supports
|
||||
- qwen3-rerank
|
||||
|
||||
(Other DashScope rerankers — gte-rerank-v2 / qwen3-vl-rerank — share the same
|
||||
endpoint but have not been validated against this transformer. Behavior with
|
||||
those models is undefined.)
|
||||
|
||||
Endpoint
|
||||
- https://dashscope.aliyuncs.com/compatible-api/v1/reranks
|
||||
|
||||
Note: chat/embed live under `/compatible-mode/v1/`, but DashScope's rerank
|
||||
route is exposed under `/compatible-api/v1/reranks` per the docs. Override
|
||||
with `DASHSCOPE_API_BASE_RERANK` to point at a different host or path.
|
||||
|
||||
Empirically, qwen3-rerank accepts `return_documents=true` and echoes
|
||||
`results[].document.text` back, even though the public docs list the flag
|
||||
as supported only for gte-rerank-v2 / qwen3-vl-rerank.
|
||||
|
||||
Docs - https://help.aliyun.com/zh/model-studio/text-rerank-api
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._uuid import uuid
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.rerank import (
|
||||
OptionalRerankParams,
|
||||
RerankBilledUnits,
|
||||
RerankResponse,
|
||||
RerankResponseMeta,
|
||||
RerankTokens,
|
||||
)
|
||||
|
||||
from ..common_utils import DashScopeError
|
||||
|
||||
DEFAULT_RERANK_URL = "https://dashscope.aliyuncs.com/compatible-api/v1/reranks"
|
||||
|
||||
|
||||
class DashScopeRerankConfig(BaseRerankConfig):
|
||||
"""
|
||||
Reference: https://help.aliyun.com/zh/model-studio/text-rerank-api
|
||||
|
||||
Targets DashScope's qwen3-rerank model. Request fields: model, query,
|
||||
documents, top_n, return_documents. Response: results[].index,
|
||||
results[].relevance_score, optionally results[].document.text (when
|
||||
return_documents=true), plus a top-level usage.total_tokens counter.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
model: str,
|
||||
optional_params: Optional[dict] = None,
|
||||
) -> str:
|
||||
if api_base is None:
|
||||
api_base = get_secret_str("DASHSCOPE_API_BASE_RERANK") or DEFAULT_RERANK_URL
|
||||
|
||||
if api_base == DEFAULT_RERANK_URL:
|
||||
return DEFAULT_RERANK_URL
|
||||
|
||||
cleaned = api_base.rstrip("/")
|
||||
if cleaned.endswith("/reranks") or cleaned.endswith("/rerank"):
|
||||
return cleaned
|
||||
|
||||
if cleaned.endswith("/v1"):
|
||||
return f"{cleaned}/reranks"
|
||||
|
||||
# Unknown base: append /reranks rather than silently ignoring the caller's api_base.
|
||||
return f"{cleaned}/reranks"
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
api_key: Optional[str] = None,
|
||||
optional_params: Optional[dict] = None,
|
||||
) -> dict:
|
||||
if api_key is None:
|
||||
api_key = get_secret_str("DASHSCOPE_API_KEY")
|
||||
if api_key is None:
|
||||
raise ValueError(
|
||||
"DashScope API key is required. Set 'DASHSCOPE_API_KEY' env var or pass api_key explicitly."
|
||||
)
|
||||
|
||||
default_headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"accept": "application/json",
|
||||
"content-type": "application/json",
|
||||
}
|
||||
return {**default_headers, **headers}
|
||||
|
||||
def get_supported_cohere_rerank_params(self, model: str) -> list:
|
||||
return ["query", "documents", "top_n", "return_documents"]
|
||||
|
||||
def map_cohere_rerank_params(
|
||||
self,
|
||||
non_default_params: Optional[dict],
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
query: str,
|
||||
documents: List[Union[str, Dict[str, Any]]],
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
top_n: Optional[int] = None,
|
||||
rank_fields: Optional[List[str]] = None,
|
||||
return_documents: Optional[bool] = True,
|
||||
max_chunks_per_doc: Optional[int] = None,
|
||||
max_tokens_per_doc: Optional[int] = None,
|
||||
) -> Dict:
|
||||
# qwen3-rerank accepts query/documents/top_n/return_documents. The
|
||||
# rest (rank_fields, max_*_per_doc) are silently dropped.
|
||||
params: OptionalRerankParams = OptionalRerankParams(
|
||||
query=query,
|
||||
documents=documents,
|
||||
)
|
||||
if top_n is not None:
|
||||
params["top_n"] = top_n
|
||||
if return_documents is not None:
|
||||
params["return_documents"] = return_documents
|
||||
return dict(params)
|
||||
|
||||
def transform_rerank_request(
|
||||
self,
|
||||
model: str,
|
||||
optional_rerank_params: Dict,
|
||||
headers: dict,
|
||||
litellm_params: Optional[dict] = None,
|
||||
) -> dict:
|
||||
if "query" not in optional_rerank_params:
|
||||
raise ValueError("query is required for DashScope rerank")
|
||||
if "documents" not in optional_rerank_params:
|
||||
raise ValueError("documents is required for DashScope rerank")
|
||||
|
||||
request: Dict[str, Any] = {
|
||||
"model": model,
|
||||
"query": optional_rerank_params["query"],
|
||||
"documents": optional_rerank_params["documents"],
|
||||
}
|
||||
if optional_rerank_params.get("top_n") is not None:
|
||||
request["top_n"] = optional_rerank_params["top_n"]
|
||||
if optional_rerank_params.get("return_documents") is not None:
|
||||
request["return_documents"] = optional_rerank_params["return_documents"]
|
||||
return request
|
||||
|
||||
def transform_rerank_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: RerankResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
api_key: Optional[str] = None,
|
||||
request_data: Optional[dict] = None,
|
||||
optional_params: Optional[dict] = None,
|
||||
litellm_params: Optional[dict] = None,
|
||||
) -> RerankResponse:
|
||||
request_data = request_data or {}
|
||||
optional_params = optional_params or {}
|
||||
litellm_params = litellm_params or {}
|
||||
try:
|
||||
response_json = raw_response.json()
|
||||
except Exception:
|
||||
raise DashScopeError(
|
||||
status_code=raw_response.status_code,
|
||||
message=raw_response.text,
|
||||
)
|
||||
|
||||
logging_obj.post_call(
|
||||
input=request_data.get("query"),
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": request_data},
|
||||
original_response=response_json,
|
||||
)
|
||||
|
||||
# DashScope error envelope: {"code": "...", "message": "...", "request_id": "..."}
|
||||
if "code" in response_json and "results" not in response_json:
|
||||
raise DashScopeError(
|
||||
status_code=raw_response.status_code,
|
||||
message=response_json.get("message", str(response_json)),
|
||||
)
|
||||
|
||||
results = response_json.get("results")
|
||||
if results is None:
|
||||
raise DashScopeError(
|
||||
status_code=raw_response.status_code,
|
||||
message=f"No results in DashScope rerank response: {response_json}",
|
||||
)
|
||||
|
||||
# qwen3-rerank returns:
|
||||
# {"index": int, "relevance_score": float}
|
||||
# plus, when return_documents=true was sent:
|
||||
# "document": {"text": "..."}
|
||||
# which already matches LiteLLM's RerankResponseDocument shape.
|
||||
transformed_results: List[dict] = []
|
||||
for r in results:
|
||||
item: Dict[str, Any] = {
|
||||
"index": r["index"],
|
||||
"relevance_score": r["relevance_score"],
|
||||
}
|
||||
doc = r.get("document")
|
||||
if isinstance(doc, dict):
|
||||
item["document"] = doc
|
||||
elif isinstance(doc, str):
|
||||
# Defensive: spec says dict, but normalize string-shaped echoes.
|
||||
item["document"] = {"text": doc}
|
||||
transformed_results.append(item)
|
||||
|
||||
usage = response_json.get("usage") or {}
|
||||
total_tokens = usage.get("total_tokens")
|
||||
billed_units = RerankBilledUnits(total_tokens=total_tokens)
|
||||
tokens = RerankTokens(input_tokens=total_tokens)
|
||||
meta = RerankResponseMeta(billed_units=billed_units, tokens=tokens)
|
||||
|
||||
return RerankResponse(
|
||||
id=response_json.get("id") or str(uuid.uuid4()),
|
||||
results=transformed_results, # type: ignore
|
||||
meta=meta,
|
||||
)
|
||||
|
||||
def get_error_class(
|
||||
self,
|
||||
error_message: str,
|
||||
status_code: int,
|
||||
headers: Union[dict, httpx.Headers],
|
||||
) -> BaseLLMException:
|
||||
if isinstance(headers, dict):
|
||||
headers = httpx.Headers(headers)
|
||||
return DashScopeError(
|
||||
status_code=status_code,
|
||||
message=error_message,
|
||||
headers=headers,
|
||||
)
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Support for OpenAI's `/v1/chat/completions` endpoint.
|
||||
Support for OpenAI's `/v1/chat/completions` endpoint.
|
||||
|
||||
Calls done in OpenAI/openai.py as DataRobot is openai-compatible.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Translate between Cohere's `/rerank` format and Deepinfra's `/rerank` format.
|
||||
Translate between Cohere's `/rerank` format and Deepinfra's `/rerank` format.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
|
|
|||
|
|
@ -2,13 +2,15 @@
|
|||
Translates from OpenAI's `/v1/chat/completions` to DeepSeek's `/v1/chat/completions`
|
||||
"""
|
||||
|
||||
from typing import Any, Coroutine, List, Literal, Optional, Tuple, Union, overload
|
||||
from typing import Any, Coroutine, List, Literal, Optional, Tuple, Union, cast, overload
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
handle_messages_with_content_list_to_str_conversion,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.utils import supports_reasoning
|
||||
|
||||
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
|
||||
|
|
@ -62,6 +64,48 @@ class DeepSeekChatConfig(OpenAIGPTConfig):
|
|||
|
||||
return optional_params
|
||||
|
||||
def _fill_reasoning_content(
|
||||
self, messages: List[AllMessageValues]
|
||||
) -> List[AllMessageValues]:
|
||||
"""
|
||||
DeepSeek thinking mode requires `reasoning_content` to be passed back on
|
||||
every assistant message in multi-turn conversations. If it is missing,
|
||||
the API returns:
|
||||
"The reasoning_content in the thinking mode must be passed back to the API."
|
||||
|
||||
For each assistant message that is missing `reasoning_content`:
|
||||
1. Promote it from `provider_specific_fields["reasoning_content"]` if present
|
||||
(LiteLLM stores provider-specific response fields there).
|
||||
2. Otherwise inject a single space — the minimum value the API accepts.
|
||||
"""
|
||||
result: List[AllMessageValues] = []
|
||||
for msg in messages:
|
||||
if msg.get("role") == "assistant" and not msg.get("reasoning_content"):
|
||||
patched = dict(cast(dict, msg))
|
||||
provider_fields = patched.get("provider_specific_fields") or {}
|
||||
stored = provider_fields.get("reasoning_content")
|
||||
if stored:
|
||||
patched["reasoning_content"] = stored
|
||||
cleaned = dict(provider_fields)
|
||||
cleaned.pop("reasoning_content", None)
|
||||
patched["provider_specific_fields"] = cleaned
|
||||
else:
|
||||
litellm.verbose_logger.warning(
|
||||
"DeepSeek thinking mode: assistant message is missing "
|
||||
"`reasoning_content` and none was saved in "
|
||||
"`provider_specific_fields`. A single-space placeholder "
|
||||
"is being injected to satisfy API validation, but the "
|
||||
"model will receive a blank reasoning chain for this turn, "
|
||||
"which may silently degrade multi-turn response quality. "
|
||||
"Preserve `reasoning_content` from the original assistant "
|
||||
"response when building multi-turn conversation history."
|
||||
)
|
||||
patched["reasoning_content"] = " "
|
||||
result.append(cast(AllMessageValues, patched))
|
||||
else:
|
||||
result.append(msg)
|
||||
return result
|
||||
|
||||
@overload
|
||||
def _transform_messages(
|
||||
self, messages: List[AllMessageValues], model: str, is_async: Literal[True]
|
||||
|
|
@ -91,6 +135,66 @@ class DeepSeekChatConfig(OpenAIGPTConfig):
|
|||
messages=messages, model=model, is_async=False
|
||||
)
|
||||
|
||||
def _thinking_mode_active(self, model: str, optional_params: dict) -> bool:
|
||||
"""
|
||||
Returns True only when thinking mode is actually active for this request:
|
||||
- model supports reasoning (capability check)
|
||||
- user explicitly passed thinking={"type": "enabled"} (opt-in check)
|
||||
"""
|
||||
return (
|
||||
supports_reasoning(model=model, custom_llm_provider="deepseek")
|
||||
and (optional_params.get("thinking") or {}).get("type") == "enabled"
|
||||
)
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Ensures `reasoning_content` is forwarded on assistant messages for
|
||||
multi-turn thinking-mode conversations (issue #28045).
|
||||
|
||||
Only runs when thinking mode is actually active - guarded by both
|
||||
supports_reasoning() (model capability) and optional_params["thinking"]
|
||||
(user explicitly enabled it), preventing spurious injection on models
|
||||
like deepseek-v3.2 that support thinking as opt-in but not always-on.
|
||||
"""
|
||||
if self._thinking_mode_active(model=model, optional_params=optional_params):
|
||||
messages = self._fill_reasoning_content(messages)
|
||||
return super().transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
async def async_transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Async equivalent of transform_request — applies the same reasoning_content
|
||||
fix for multi-turn thinking-mode conversations.
|
||||
"""
|
||||
if self._thinking_mode_active(model=model, optional_params=optional_params):
|
||||
messages = self._fill_reasoning_content(messages)
|
||||
return await super().async_transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
def _get_openai_compatible_provider_info(
|
||||
self, api_base: Optional[str], api_key: Optional[str]
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue