Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_feat/v1.84.0-mcp-gateway-jwt-auth

# Conflicts:
#	litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py
This commit is contained in:
mateo-berri 2026-05-26 18:41:20 +00:00
commit ccea19d66f
No known key found for this signature in database
809 changed files with 34347 additions and 5189 deletions

View file

@ -2477,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
@ -2497,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
@ -2533,7 +2541,6 @@ jobs:
paths:
- litellm-docker-database.tar.zst
test_bad_database_url:
machine:
image: ubuntu-2204:2024.04.1

View file

@ -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 }}"

View file

@ -0,0 +1,47 @@
name: Create Daily oss-agent-shin Branch
on:
schedule:
- cron: "0 0 * * *" # Runs every day at midnight UTC
workflow_dispatch: # Allow manual trigger
jobs:
create-oss-agent-shin-branch:
if: github.repository == 'BerriAI/litellm'
runs-on: ubuntu-latest
permissions:
contents: write
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
fetch-depth: 0
persist-credentials: false
- name: Create daily oss-agent-shin branch
env:
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
run: |
# Configure Git user
git config user.name "github-actions[bot]"
git config user.email "github-actions[bot]@users.noreply.github.com"
# Generate branch name with MM_DD_YYYY format
BRANCH_NAME="litellm_oss_agent_shin_$(date +'%m_%d_%Y')"
echo "Creating branch: $BRANCH_NAME"
# Fetch all branches
git fetch --all
# Check if the branch already exists
if git show-ref --verify --quiet refs/remotes/origin/$BRANCH_NAME; then
echo "Branch $BRANCH_NAME already exists. Skipping creation."
else
echo "Creating new branch: $BRANCH_NAME"
# Create the new branch from main
git checkout -b $BRANCH_NAME origin/main
# Push the new branch
git push origin $BRANCH_NAME
echo "Successfully created and pushed branch: $BRANCH_NAME"
fi

View file

@ -7,6 +7,7 @@ on:
- litellm_internal_staging
- litellm_oss_branch
- "litellm_**"
workflow_dispatch:
permissions:
contents: read
@ -42,3 +43,16 @@ jobs:
workers: 2
reruns: 2
artifact-name: proxy-endpoints
# Behavior-pinning tests for litellm/proxy/proxy_server.py. Owns its
# own job (not a path on the proxy-endpoints job above) so its budget
# is independent and its coverage artifact is uploaded separately.
# See: https://www.notion.so/36c43b8acdab81ee845fd5365128a2fc
proxy-server:
uses: ./.github/workflows/_test-unit-base.yml
with:
test-path: tests/test_litellm/proxy/proxy_server
workers: 4
reruns: 2
timeout-minutes: 60
artifact-name: proxy-server

View file

@ -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"

View file

@ -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

View file

@ -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 }}

View file

@ -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: []

View file

@ -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 \
@ -54,22 +55,10 @@ COPY . .
# Set non-root flag for build time consistency
ENV LITELLM_NON_ROOT=true
# Stage the pre-built Admin UI from the checked-in Next.js static export.
# _experimental/out/ is regenerated as part of the release runbook.
# Restructure extensionless routes (foo.html -> foo/index.html) to match the layout
# proxy_server.py expects, and drop a readiness marker.
RUN mkdir -p /var/lib/litellm/ui /var/lib/litellm/assets && \
cp -r /app/litellm/proxy/_experimental/out/. /var/lib/litellm/ui/ && \
cp /app/litellm/proxy/logo.jpg /var/lib/litellm/assets/logo.jpg && \
( cd /var/lib/litellm/ui && \
for html_file in *.html; do \
if [ "$html_file" != "index.html" ] && [ -f "$html_file" ]; then \
folder_name="${html_file%.html}" && \
mkdir -p "$folder_name" && \
mv "$html_file" "$folder_name/index.html"; \
fi; \
done && \
touch .litellm_ui_ready )
touch /var/lib/litellm/ui/.litellm_ui_ready
RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
if [ "$PROXY_EXTRAS_SOURCE" = "published" ]; then \

File diff suppressed because one or more lines are too long

View file

@ -1877,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,
)

View file

@ -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",

View file

@ -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

View file

@ -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] = {

View file

@ -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[

View file

@ -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"
)

View file

@ -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"

View file

@ -24,6 +24,7 @@ from litellm.litellm_core_utils.llm_cost_calc.usage_object_transformation import
from litellm.litellm_core_utils.llm_cost_calc.utils import (
CostCalculatorUtils,
_generic_cost_per_character,
_get_regional_uplift_multiplier,
_get_service_tier_cost_key,
_parse_prompt_tokens_details,
calculate_cost_component,
@ -312,6 +313,10 @@ def cost_per_token( # noqa: PLR0915
audio_transcription_file_duration: float = 0.0, # for audio transcription calls - the file time in seconds
### SERVICE TIER ###
service_tier: Optional[str] = None, # for OpenAI service tier pricing
### DATA RESIDENCY ###
data_residency: Optional[
str
] = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
response: Optional[Any] = None,
### REQUEST MODEL ###
request_model: Optional[str] = None, # original request model for router detection
@ -493,6 +498,7 @@ def cost_per_token( # noqa: PLR0915
usage=usage_block,
custom_llm_provider=custom_llm_provider,
service_tier=service_tier,
data_residency=data_residency,
)
return prompt_cost, completion_cost
@ -521,7 +527,10 @@ def cost_per_token( # noqa: PLR0915
or call_type == CallTypes.retrieve_batch
):
return batch_cost_calculator(
usage=usage_block, model=model, custom_llm_provider=custom_llm_provider
usage=usage_block,
model=model,
custom_llm_provider=custom_llm_provider,
data_residency=data_residency,
)
elif call_type == "atranscription" or call_type == "transcription":
if _transcription_usage_has_token_details(usage_block):
@ -529,6 +538,7 @@ def cost_per_token( # noqa: PLR0915
model=model_without_prefix,
usage=usage_block,
service_tier=service_tier,
data_residency=data_residency,
)
return openai_cost_per_second(
@ -579,7 +589,10 @@ def cost_per_token( # noqa: PLR0915
)
elif custom_llm_provider == "openai":
return openai_cost_per_token(
model=model, usage=usage_block, service_tier=service_tier
model=model,
usage=usage_block,
service_tier=service_tier,
data_residency=data_residency,
)
elif custom_llm_provider == "databricks":
return databricks_cost_per_token(model=model, usage=usage_block)
@ -631,6 +644,7 @@ def cost_per_token( # noqa: PLR0915
usage=usage_block,
custom_llm_provider=custom_llm_provider,
service_tier=service_tier,
data_residency=data_residency,
)
if (
@ -1117,6 +1131,10 @@ def completion_cost( # noqa: PLR0915
litellm_logging_obj: Optional[LitellmLoggingObject] = None,
### SERVICE TIER ###
service_tier: Optional[str] = None, # for OpenAI service tier pricing
### DATA RESIDENCY ###
data_residency: Optional[
str
] = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
) -> float:
"""
Calculate the cost of a given completion call fot GPT-3.5-turbo, llama2, any litellm supported llm.
@ -1516,6 +1534,7 @@ def completion_cost( # noqa: PLR0915
combined_usage_object=cost_per_token_usage_object,
custom_llm_provider=custom_llm_provider,
litellm_model_name=model,
data_residency=data_residency,
)
elif call_type == _MCP_CALL_TYPE:
from litellm.proxy._experimental.mcp_server.cost_calculator import (
@ -1600,6 +1619,7 @@ def completion_cost( # noqa: PLR0915
audio_transcription_file_duration=audio_transcription_file_duration,
rerank_billed_units=rerank_billed_units,
service_tier=service_tier,
data_residency=data_residency,
response=completion_response,
request_model=request_model_for_cost,
)
@ -1811,6 +1831,10 @@ def response_cost_calculator(
litellm_logging_obj: Optional[LitellmLoggingObject] = None,
### SERVICE TIER ###
service_tier: Optional[str] = None, # for OpenAI service tier pricing
### DATA RESIDENCY ###
data_residency: Optional[
str
] = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
) -> float:
"""
Returns
@ -1844,6 +1868,7 @@ def response_cost_calculator(
router_model_id=router_model_id,
litellm_logging_obj=litellm_logging_obj,
service_tier=service_tier,
data_residency=data_residency,
)
return response_cost
except Exception as e:
@ -2202,6 +2227,7 @@ def batch_cost_calculator(
model: str,
custom_llm_provider: Optional[str] = None,
model_info: Optional[ModelInfo] = None,
data_residency: Optional[str] = None,
) -> Tuple[float, float]:
"""
Calculate the cost of a batch job.
@ -2286,6 +2312,11 @@ def batch_cost_calculator(
usage.completion_tokens * (output_cost_per_token) / 2
) # batch cost is usually half of the regular token cost
uplift = _get_regional_uplift_multiplier(model_info, data_residency)
if uplift != 1.0:
total_prompt_cost *= uplift
total_completion_cost *= uplift
return total_prompt_cost, total_completion_cost
@ -2431,6 +2462,7 @@ def handle_realtime_stream_cost_calculation(
combined_usage_object: Usage,
custom_llm_provider: str,
litellm_model_name: str,
data_residency: Optional[str] = None,
) -> float:
"""
Handles the cost calculation for realtime stream responses.
@ -2461,6 +2493,7 @@ def handle_realtime_stream_cost_calculation(
model=model_name,
usage=combined_usage_object,
custom_llm_provider=custom_llm_provider,
data_residency=data_residency,
)
except Exception:
continue

View file

@ -1,6 +1,5 @@
from typing import AsyncIterator, Dict, Iterator, Literal, NamedTuple, Union
FileContentProvider = Literal[
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic", "manus"
]

View file

@ -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)
"""

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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.
"""

View file

@ -9,7 +9,6 @@ import polars as pl
from .schema import FOCUS_NORMALIZED_SCHEMA
_TAG_KEYS = (
"team_id",
"team_alias",

View file

@ -702,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)
@ -726,9 +734,57 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
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,
@ -750,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)
@ -937,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."""
@ -1023,8 +1082,10 @@ 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.
@ -1617,6 +1678,37 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
"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_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))
@ -2962,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,
@ -3012,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,

View file

@ -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.

View file

@ -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,

View file

@ -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
"""

View file

@ -49,7 +49,6 @@ from litellm.types.interactions import InteractionEnvironment
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import client
# ------------------------------------------------------------------ #
# Shared helpers #
# ------------------------------------------------------------------ #

View file

@ -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="...")
"""

View file

@ -1,5 +1,7 @@
from typing import Optional
from litellm.llms.openai.data_residency import infer_openai_data_residency
# Pre-define optional kwargs keys as frozenset for O(1) lookups
# These are extracted from kwargs only if present, avoiding unnecessary .get() calls
_OPTIONAL_KWARGS_KEYS = frozenset(
@ -103,6 +105,10 @@ def get_litellm_params(
if litellm_trace_id is None:
litellm_trace_id = _meta.get("trace_id") or _meta.get("session_id")
data_residency: Optional[str] = infer_openai_data_residency(
custom_llm_provider, api_base
)
# Build base dict with explicit parameters (always included)
litellm_params = {
"acompletion": acompletion,
@ -112,6 +118,7 @@ def get_litellm_params(
"verbose": verbose,
"custom_llm_provider": custom_llm_provider,
"api_base": api_base,
"data_residency": data_residency,
"litellm_call_id": litellm_call_id,
"model_alias_map": model_alias_map,
"completion_call_id": completion_call_id,

View file

@ -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":

View file

@ -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(
@ -1552,6 +1546,11 @@ class Logging(LiteLLMLoggingBaseClass):
if self.optional_params
else None
),
"data_residency": (
self.litellm_params.get("data_residency")
if hasattr(self, "litellm_params") and self.litellm_params
else None
),
}
except Exception as e: # error creating kwargs for cost calculation
debug_info = StandardLoggingModelCostFailureDebugInformation(
@ -5146,13 +5145,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 ""

View file

@ -9,6 +9,7 @@ from litellm.types.utils import (
CacheCreationTokenDetails,
CallTypes,
CompletionTokensDetailsWrapper,
DataResidency,
ImageResponse,
ModelInfo,
PassthroughCallTypes,
@ -617,11 +618,46 @@ def _calculate_input_cost(
return prompt_cost
def _get_regional_uplift_multiplier(
model_info: ModelInfo, data_residency: Optional[str]
) -> float:
"""
Resolve the per-model regional-processing uplift multiplier for a given
data-residency region.
OpenAI applies a flat percentage uplift (e.g. +10%) on all token costs for
requests served from a regionalized hostname (eu./us.api.openai.com). The
multiplier is stored on the model entry as
``regional_processing_uplift_multiplier_<region>`` (e.g. 1.10).
Returns 1.0 (no uplift) when ``data_residency`` is ``None`` or when the
model has no multiplier configured for the given region.
"""
if data_residency is None:
return 1.0
residency = data_residency.lower()
if residency not in {r.value for r in DataResidency}:
return 1.0
multiplier = model_info.get(f"regional_processing_uplift_multiplier_{residency}")
if multiplier is None:
return 1.0
try:
return float(cast(float, multiplier))
except (TypeError, ValueError):
verbose_logger.exception(
"Invalid regional_processing_uplift_multiplier_%s for model; "
"defaulting to 1.0",
residency,
)
return 1.0
def generic_cost_per_token( # noqa: PLR0915
model: str,
usage: Usage,
custom_llm_provider: str,
service_tier: Optional[str] = None,
data_residency: Optional[str] = None,
) -> Tuple[float, float]:
"""
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
@ -631,6 +667,8 @@ def generic_cost_per_token( # noqa: PLR0915
Input:
- model: str, the model name without provider prefix
- usage: LiteLLM Usage block, containing anthropic caching information
- data_residency: optional OpenAI data-residency region (e.g. "eu", "us"),
used to apply the per-model regional-processing uplift multiplier.
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
@ -781,6 +819,14 @@ def generic_cost_per_token( # noqa: PLR0915
)
completion_cost += float(image_tokens) * _output_cost_per_image_token
## REGIONAL DATA-RESIDENCY UPLIFT
# Applied as a flat multiplier across all token costs for the request
# when the upstream is a regionalized OpenAI host (eu./us.api.openai.com).
uplift = _get_regional_uplift_multiplier(model_info, data_residency)
if uplift != 1.0:
prompt_cost *= uplift
completion_cost *= uplift
return prompt_cost, completion_cost

View file

@ -1204,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

View file

@ -5590,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

View file

@ -146,6 +146,37 @@ class SensitiveDataMasker:
return masked_data
_default_masker = SensitiveDataMasker()
def mask_sensitive_keys(
data: Dict[str, Any], sensitive_fields: Set[str]
) -> Dict[str, Any]:
"""Return a new dict with values masked for keys listed in ``sensitive_fields``.
Unlike :meth:`SensitiveDataMasker.mask_dict`, this does exact key-name
matching (not segment matching), so callers explicitly enumerate which
fields to mask. Non-string and None values are passed through unchanged.
Values shorter than ``visible_prefix + visible_suffix`` (8 by default)
fall outside :meth:`SensitiveDataMasker._mask_value`'s partial-reveal
range and are replaced with a fixed-length all-mask string, so a short
credential is never returned verbatim.
"""
masked: Dict[str, Any] = {}
mask_char = _default_masker.mask_char
min_visible = _default_masker.visible_prefix + _default_masker.visible_suffix
for key, value in data.items():
if value is not None and key in sensitive_fields and isinstance(value, str):
if len(value) < min_visible:
masked[key] = mask_char * len(value) if value else value
else:
masked[key] = _default_masker._mask_value(value)
else:
masked[key] = value
return masked
# Usage example:
"""
masker = SensitiveDataMasker()

View file

@ -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.
"""

View file

@ -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 (

View file

@ -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)
# ---------------------------------------------------------------------------

View file

@ -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)

View file

@ -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.

View file

@ -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
}

View file

@ -0,0 +1,3 @@
from .transformation import AzureSpeechAudioTranscriptionConfig
__all__ = ["AzureSpeechAudioTranscriptionConfig"]

View 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 ""

View file

@ -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,

View file

@ -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)
"""

View file

@ -1,3 +1,5 @@
import asyncio
import hashlib
import json
import os
from typing import Any, Callable, Dict, Literal, NamedTuple, Optional, Union, cast
@ -449,6 +451,25 @@ class BaseAzureLLM(BaseOpenAILLM):
] = None
client_initialization_params: dict = locals()
client_initialization_params["is_async"] = _is_async
_lp = litellm_params or {}
_ad_provider = _lp.get("azure_ad_token_provider")
_ad_token = _lp.get("azure_ad_token")
_client_secret = _lp.get("client_secret")
_azure_password = _lp.get("azure_password")
client_initialization_params["azure_ad_token"] = (
hashlib.sha256(_ad_token.encode()).hexdigest()
if isinstance(_ad_token, str)
else None
)
client_initialization_params["azure_ad_token_provider"] = (
f"provider_id={id(_ad_provider) if callable(_ad_provider) else None}"
f"|tenant_id={_lp.get('tenant_id')}"
f"|client_id={_lp.get('client_id')}"
f"|client_secret={hashlib.sha256(_client_secret.encode()).hexdigest() if isinstance(_client_secret, str) else None}"
f"|azure_username={_lp.get('azure_username')}"
f"|azure_password={hashlib.sha256(_azure_password.encode()).hexdigest() if isinstance(_azure_password, str) else None}"
f"|azure_scope={_lp.get('azure_scope')}"
)
if client is None:
cached_client = self.get_cached_openai_client(
client_initialization_params=client_initialization_params,
@ -474,8 +495,29 @@ class BaseAzureLLM(BaseOpenAILLM):
if self._is_azure_v1_api_version(api_version):
# Extract only params that OpenAI client accepts
# Always use /openai/v1/ regardless of whether user passed "v1", "latest", or "preview"
v1_params = {
"api_key": azure_client_params.get("api_key"),
# The OpenAI client accepts a callable for `api_key` and re-invokes it
# on every request (via `_refresh_api_key`), so passing
# `azure_ad_token_provider` directly preserves Azure AD token refresh
# behavior that the regular AzureOpenAI client provides.
v1_api_key: Optional[Union[str, Callable[[], Any]]] = (
azure_client_params.get("api_key")
or azure_client_params.get("azure_ad_token_provider")
or azure_client_params.get("azure_ad_token")
)
if _is_async is True and callable(v1_api_key):
# AsyncOpenAI expects an async provider; wrap the sync provider
# returned by azure-identity. Offload to a thread so a token
# refresh (blocking HTTP call to AAD on cache miss) does not
# stall the event loop.
_sync_provider = v1_api_key
async def _async_v1_api_key() -> str:
return await asyncio.to_thread(_sync_provider)
v1_api_key = _async_v1_api_key
v1_params: Dict[str, Any] = {
"api_key": v1_api_key,
"base_url": f"{api_base}/openai/v1/",
}
if "timeout" in azure_client_params:

View file

@ -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

View file

@ -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

View file

@ -177,8 +177,14 @@ def extract_model_id_from_unified_id(
if decoded_id:
unified_id = decoded_id
# Extract model ID
match = re.search(r"model_id,([^;]+)", unified_id)
# Extract model ID. Anchor to a field boundary (start of string or
# after `;`) so this regex doesn't substring-match the `model_id,`
# inside file_id encodings' `llm_output_file_model_id,<deployment_uuid>`
# field — that would feed the deployment UUID as a model candidate
# into the team-access check and 403 every team-BYOK file attach
# with `Tried to access <uuid>` (LIT-3244 patch/1.86.0 second-order
# finding).
match = re.search(r"(?:^|;)model_id,([^;]+)", unified_id)
if match:
return match.group(1).strip()

View file

@ -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)

View file

@ -157,8 +157,8 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
def _get_agent_runtime_arn(self, model: str) -> str:
"""
Extract ARN from model string
model = "agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC"
returns: "arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC"
model = "agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp"
returns: "arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp"
"""
parts = model.split("/", 1)
if len(parts) != 2 or parts[0] != "agentcore":
@ -170,7 +170,7 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
def _extract_region_from_arn(self, arn: str) -> str:
"""
Extract region from ARN
arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC
arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp
returns: us-west-2
"""
parts = arn.split(":")

View file

@ -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"
)

View file

@ -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

View file

@ -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
"""

View file

@ -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"

View file

@ -1,5 +1,5 @@
"""
Legacy /v1/embedding handler for Bedrock Cohere.
Legacy /v1/embedding handler for Bedrock Cohere.
"""
import json

View file

@ -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,

View file

@ -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(
@ -1856,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,
@ -2047,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,
@ -2061,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,
@ -2083,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,
)
@ -4590,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,

View file

@ -13,7 +13,6 @@ from typing import Tuple
import httpx
# ---------------------------------------------------------------------------
# Pre-built response templates
# ---------------------------------------------------------------------------

View file

@ -1,5 +1,5 @@
"""
Cost calculator for Dashscope Chat models.
Cost calculator for Dashscope Chat models.
Handles tiered pricing and prompt caching scenarios.
"""

View file

@ -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.
"""

View file

@ -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

View file

@ -1,5 +1,5 @@
"""
Cost calculator for DeepSeek Chat models.
Cost calculator for DeepSeek Chat models.
Handles prompt caching scenario.
"""

View file

@ -22,7 +22,6 @@ from litellm.types.utils import all_litellm_params
from ..common_utils import ElevenLabsException
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.llms.openai import HttpxBinaryResponseContent

View file

@ -23,7 +23,6 @@ from litellm.types.agents import (
AgentVersionsResponse,
)
# Keys inside litellm_params that should be forwarded to the Gemini
# create-agent body verbatim.
_GEMINI_AGENT_BODY_KEYS = ("base_agent", "instructions", "base_environment")

View file

@ -55,7 +55,7 @@ def _convert_image_to_gemini_format(image_file) -> Dict[str, str]:
def _usage_video_resolution_from_parameters(
parameters: Dict[str, Any]
parameters: Dict[str, Any],
) -> Optional[str]:
"""Normalize Veo ``parameters.resolution`` for usage and cost tracking."""
res = parameters.get("resolution")

View file

@ -1,5 +1,5 @@
"""
Transformation logic from Cohere's /v1/rerank format to Infinity's `/v1/rerank` format.
Transformation logic from Cohere's /v1/rerank format to Infinity's `/v1/rerank` format.
Why separate file? Make it easy to see how transformation works
"""

View file

@ -1,5 +1,5 @@
"""
Transformation logic from Cohere's /v1/rerank format to Jina AI's `/v1/rerank` format.
Transformation logic from Cohere's /v1/rerank format to Jina AI's `/v1/rerank` format.
Why separate file? Make it easy to see how transformation works

View file

@ -1,5 +1,5 @@
"""
Transformation logic from OpenAI /v1/embeddings format to LM Studio's `/v1/embeddings` format.
Transformation logic from OpenAI /v1/embeddings format to LM Studio's `/v1/embeddings` format.
Why separate file? Make it easy to see how transformation works

View file

@ -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 Novita AI is openai-compatible.

View file

@ -1,7 +1,7 @@
"""
Nvidia NIM endpoint: https://docs.api.nvidia.com/nim/reference/databricks-dbrx-instruct-infer
Nvidia NIM endpoint: https://docs.api.nvidia.com/nim/reference/databricks-dbrx-instruct-infer
This is OpenAI compatible
This is OpenAI compatible
This file only contains param mapping logic

View file

@ -1,7 +1,7 @@
"""
Nvidia NIM embeddings endpoint: https://docs.api.nvidia.com/nim/reference/nvidia-nv-embedqa-e5-v5-infer
This is OpenAI compatible
This is OpenAI compatible
This file only contains param mapping logic

View file

@ -0,0 +1,386 @@
"""
OCI Generative AI — Cohere-specific chat transformation helpers.
Handles message history building, tool definition adaptation, non-streaming
response parsing, and streaming chunk parsing for models served with
``apiFormat="COHERE"`` (e.g. ``cohere.command-*``).
"""
import datetime
import json
from typing import Any, Dict, List, Optional
import httpx
from pydantic import ValidationError
from litellm.llms.oci.chat.generic import (
_normalize_oci_finish_reason,
_synthesize_oci_tool_call_id,
)
from litellm.llms.oci.common_utils import (
OCI_JSON_TO_PYTHON_TYPES,
OCIError,
enrich_cohere_param_description,
resolve_oci_schema_anyof,
resolve_oci_schema_refs,
sanitize_oci_schema,
)
from litellm.types.llms.oci import (
CohereChatResult,
CohereMessage,
CohereParameterDefinition,
CohereStreamChunk,
CohereTool,
CohereToolCall,
CohereToolMessage,
CohereToolResult,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import (
Choices,
Delta,
ModelResponse,
ModelResponseStream,
StreamingChoices,
)
from litellm.types.utils import Usage
def _extract_text_content(content: Any) -> str:
"""Return the plain-text representation of a message content value."""
if content is None:
return ""
if isinstance(content, str):
return content
if isinstance(content, list):
return "".join(
item.get("text", "")
for item in content
if isinstance(item, dict) and item.get("type") == "text"
)
return str(content)
def adapt_messages_to_cohere_standard(
messages: List[AllMessageValues],
) -> List[CohereMessage]:
"""Build a Cohere ``chatHistory`` list from an OpenAI-format message array.
- All messages except the *last user message* are included. The caller pulls
the last user message into the request's top-level ``message`` field, so
trailing tool results (the standard agentic continuation pattern) still
appear in ``chatHistory`` and reach the model.
- If no user message exists, every message is included (no slice).
- System messages must be filtered out by the caller (they are routed into
``preambleOverride`` separately) — they are not represented in
``chatHistory``.
- Tool results are expressed as OCI ``CohereToolMessage.toolResults`` entries,
with the originating call's name and parameters resolved from the preceding
assistant message via a ``tool_call_id`` lookup.
"""
# First pass: build tool_call_id → CohereToolCall so tool-result messages can
# reference the originating call by name and parameters.
tool_call_lookup: Dict[str, CohereToolCall] = {}
for msg in messages:
if msg.get("role") == "assistant":
tool_calls_raw: Any = msg.get("tool_calls") or []
for tc in tool_calls_raw:
tc_id = tc.get("id", "")
raw_args: Any = tc.get("function", {}).get("arguments", "{}")
try:
params: Dict[str, Any] = (
json.loads(raw_args) if isinstance(raw_args, str) else raw_args
)
except json.JSONDecodeError:
params = {}
tool_call_lookup[tc_id] = CohereToolCall(
name=str(tc.get("function", {}).get("name", "")),
parameters=params,
)
last_user_index = next(
(
i
for i in range(len(messages) - 1, -1, -1)
if messages[i].get("role") == "user"
),
None,
)
history_source = (
messages
if last_user_index is None
else [m for i, m in enumerate(messages) if i != last_user_index]
)
chat_history: List[CohereMessage] = []
for msg in history_source:
role = msg.get("role")
content = _extract_text_content(msg.get("content"))
tool_calls: Optional[List[CohereToolCall]] = None
if role == "assistant" and msg.get("tool_calls"): # type: ignore[union-attr,typeddict-item]
tool_calls = []
for tc in msg["tool_calls"]: # type: ignore[union-attr,typeddict-item]
raw_arguments: Any = tc.get("function", {}).get("arguments", {})
if isinstance(raw_arguments, str):
try:
arguments: Dict[str, Any] = json.loads(raw_arguments)
except json.JSONDecodeError:
arguments = {}
else:
arguments = raw_arguments
tool_calls.append(
CohereToolCall(
name=str(tc.get("function", {}).get("name", "")),
parameters=arguments,
)
)
if role == "user":
chat_history.append(CohereMessage(role="USER", message=content))
elif role == "assistant":
chat_history.append(
CohereMessage(role="CHATBOT", message=content, toolCalls=tool_calls)
)
elif role == "tool":
tool_call_id = str(msg.get("tool_call_id", "") or "")
cohere_call = tool_call_lookup.get(
tool_call_id, CohereToolCall(name="", parameters={})
)
tool_result = CohereToolResult(
call=cohere_call,
outputs=[{"output": content}],
)
# OpenAI emits one tool-role message per parallel tool call, but
# the OCI Cohere API expects all results from a single assistant
# turn to share one TOOL history entry with multiple toolResults.
# Merge consecutive tool messages so the model sees the parallel
# call/result pairing correctly during agentic loops.
if chat_history and isinstance(chat_history[-1], CohereToolMessage):
chat_history[-1].toolResults.append(tool_result)
else:
chat_history.append(CohereToolMessage(toolResults=[tool_result]))
return chat_history
def adapt_tool_definitions_to_cohere_standard(
tools: List[Dict[str, Any]],
) -> List[CohereTool]:
"""Adapt OpenAI-format tool definitions to the OCI Cohere format.
- Resolves ``$ref``/``$defs`` and ``anyOf`` patterns that OCI rejects.
- Maps JSON Schema type names to Python type names (``"string"`` → ``"str"``).
- Embeds unsupported constraints (enum, format, range, pattern) into the
parameter description so the model can still see them.
"""
cohere_tools = []
for tool in tools:
function_def = tool.get("function", {})
raw_params = function_def.get("parameters", {})
resolved = sanitize_oci_schema(
resolve_oci_schema_anyof(resolve_oci_schema_refs(raw_params))
)
properties = resolved.get("properties", {})
required = resolved.get("required", [])
parameter_definitions = {}
for param_name, param_schema in properties.items():
json_type = param_schema.get("type", "string")
python_type = OCI_JSON_TO_PYTHON_TYPES.get(json_type, json_type)
parameter_definitions[param_name] = CohereParameterDefinition(
description=enrich_cohere_param_description(
param_schema.get("description", ""), param_schema
),
type=python_type,
isRequired=param_name in required,
)
cohere_tools.append(
CohereTool(
name=function_def.get("name", ""),
description=function_def.get("description", ""),
parameterDefinitions=parameter_definitions,
)
)
return cohere_tools
def handle_cohere_response(
json_response: dict,
model: str,
model_response: ModelResponse,
raw_response: httpx.Response,
) -> ModelResponse:
"""Parse a non-streaming Cohere OCI response into a LiteLLM ModelResponse."""
try:
cohere_response = CohereChatResult(**json_response)
except (TypeError, ValidationError) as e:
raise OCIError(
message=f"Response cannot be casted to CohereChatResult: {str(e)}",
status_code=raw_response.status_code,
)
model_response.model = model
model_response.created = int(datetime.datetime.now().timestamp())
response_text = cohere_response.chatResponse.text
finish_reason = _normalize_oci_finish_reason(
cohere_response.chatResponse.finishReason
)
tool_calls: Optional[List[Dict[str, Any]]] = None
if cohere_response.chatResponse.toolCalls:
tool_calls = [
{
"id": _synthesize_oci_tool_call_id(
i, tc.name, json.dumps(tc.parameters, sort_keys=True)
),
"type": "function",
"function": {
"name": tc.name,
"arguments": json.dumps(tc.parameters),
},
}
for i, tc in enumerate(cohere_response.chatResponse.toolCalls)
]
content: Optional[str] = response_text if response_text else None
# Only include ``tool_calls`` in the message dict when actually present.
# Passing an explicit ``None`` would let downstream consumers that key off
# ``"tool_calls" in message`` (rather than truthiness) incorrectly conclude
# that tool calls were attempted. Matches the generic handler's behaviour,
# which only sets ``message.tool_calls`` when tool calls are present.
message: Dict[str, Any] = {"role": "assistant", "content": content}
if tool_calls is not None:
message["tool_calls"] = tool_calls
model_response.choices = [
Choices(
index=0,
message=message,
finish_reason=finish_reason,
)
]
usage_info = cohere_response.chatResponse.usage
if usage_info is not None:
model_response.usage = Usage( # type: ignore[attr-defined]
prompt_tokens=usage_info.promptTokens,
completion_tokens=usage_info.completionTokens,
total_tokens=usage_info.totalTokens,
)
else:
model_response.usage = Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0) # type: ignore[attr-defined]
return model_response
def handle_cohere_stream_chunk(
dict_chunk: dict,
prior_tool_calls_emitted: bool = False,
prior_text_emitted: bool = False,
) -> ModelResponseStream:
"""Parse a single Cohere SSE chunk into a LiteLLM ModelResponseStream.
``prior_tool_calls_emitted`` lets the caller signal whether tool calls
were already emitted in earlier chunks of the same stream. When set, the
terminal consolidation chunk's tool calls are suppressed (they would
duplicate prior deltas); otherwise they are passed through so a stream
that delivers tool calls only on the terminal chunk doesn't silently
drop them.
``prior_text_emitted`` plays the analogous role for the ``text`` field:
when set, the terminal consolidation chunk's ``text`` is suppressed
(it would re-emit the full assembled response on top of prior deltas);
when unset (e.g. a degenerate stream that delivers the entire response
in a single SSE event carrying both ``chatHistory`` and ``finishReason``),
the text is passed through so the response content isn't silently lost.
"""
try:
typed_chunk = CohereStreamChunk(**dict_chunk)
except (TypeError, ValidationError) as e:
raise OCIError(
status_code=500,
message=f"Chunk cannot be parsed as CohereStreamChunk: {str(e)}",
)
if typed_chunk.index is None:
typed_chunk.index = 0
# OCI Cohere's terminal SSE event re-sends the full assembled response in
# `text` alongside a populated `chatHistory` and a non-null `finishReason`.
# Emitting that text would concatenate the whole response onto the
# already-streamed deltas. We require both signals to be present so that a
# future API change which adds `chatHistory` to intermediate chunks (or a
# rare early-populated case) doesn't silently drop legitimate token deltas.
is_terminal_consolidation = (
typed_chunk.chatHistory is not None and typed_chunk.finishReason is not None
)
# On non-terminal text-free chunks (e.g. tool-call-only or keep-alive
# chunks) emit ``content=None`` rather than ``content=""`` so downstream
# stream-mergers that distinguish "no text in this delta" from "an
# explicitly empty text delta" behave correctly.
#
# We only suppress the terminal chunk's ``text`` when the caller has
# confirmed that text deltas were already emitted earlier — otherwise
# (e.g. a degenerate stream that delivers the whole response in a
# single SSE event), passing it through is the only chance to surface it.
text: Optional[str] = (
None if (is_terminal_consolidation and prior_text_emitted) else typed_chunk.text
)
# Tool calls on the terminal consolidation chunk (whether from
# `typed_chunk.toolCalls` or from `chatHistory`) typically restate what
# was already streamed in intermediate chunks. Re-emitting them would
# mint fresh `uuid4` IDs and cause downstream consumers to execute each
# tool call twice. We only suppress when the caller has confirmed that
# tool calls were already emitted earlier — otherwise (e.g. a short
# response that delivers tool calls exclusively on the terminal chunk),
# passing them through is the only chance to surface them.
cohere_tool_calls = (
None
if (is_terminal_consolidation and prior_tool_calls_emitted)
else typed_chunk.toolCalls
)
tool_calls: Optional[List[Dict[str, Any]]] = None
if cohere_tool_calls:
tool_calls = [
{
# Cohere protocol has no tool-call id, so we synthesize one
# deterministically from the call's content/position. A random
# uuid4 per chunk would cause downstream stream-mergers to
# treat each chunk as a distinct tool call.
"id": _synthesize_oci_tool_call_id(
i, tc.name, json.dumps(tc.parameters, sort_keys=True)
),
"type": "function",
"function": {
"name": tc.name,
"arguments": json.dumps(tc.parameters),
},
}
for i, tc in enumerate(cohere_tool_calls)
]
finish_reason = _normalize_oci_finish_reason(typed_chunk.finishReason)
return ModelResponseStream(
choices=[
StreamingChoices(
index=typed_chunk.index,
delta=Delta(
content=text,
tool_calls=tool_calls,
provider_specific_fields=None,
thinking_blocks=None,
reasoning_content=None,
),
finish_reason=finish_reason,
)
]
)

View file

@ -0,0 +1,477 @@
"""
OCI Generative AI — Generic-format chat transformation helpers.
Handles message building, tool definition adaptation, non-streaming response
parsing, and streaming chunk parsing for models served with
``apiFormat="GENERIC"`` (e.g. Meta Llama, xAI Grok, Google Gemini).
"""
import datetime
import hashlib
from typing import Any, Dict, List, Optional, Union
import httpx
from pydantic import ValidationError
from litellm.llms.oci.common_utils import (
OCIError,
resolve_oci_schema_anyof,
resolve_oci_schema_refs,
sanitize_oci_schema,
)
from litellm.types.llms.oci import (
OCICompletionResponse,
OCIContentPartUnion,
OCIImageContentPart,
OCIImageUrl,
OCIMessage,
OCIRoles,
OCIStreamChunk,
OCITextContentPart,
OCIToolCall,
OCIToolDefinition,
OCIVendors,
)
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import (
Delta,
ModelResponse,
ModelResponseStream,
StreamingChoices,
)
from litellm.types.utils import ChatCompletionMessageToolCall, Usage
# Maps OpenAI role names to OCI GENERIC role names.
open_ai_to_generic_oci_role_map: Dict[str, OCIRoles] = {
"system": "SYSTEM",
"user": "USER",
"assistant": "ASSISTANT",
"tool": "TOOL",
}
# ---------------------------------------------------------------------------
# Message building
# ---------------------------------------------------------------------------
def adapt_messages_to_generic_oci_standard_content_message(
role: str, content: Union[str, list]
) -> OCIMessage:
"""Convert a plain-text or multipart content message to OCI format."""
new_content: List[OCIContentPartUnion] = []
if isinstance(content, str):
return OCIMessage(
role=open_ai_to_generic_oci_role_map[role],
content=[OCITextContentPart(text=content)],
toolCalls=None,
toolCallId=None,
)
for content_item in content:
if not isinstance(content_item, dict):
raise OCIError(
status_code=400, message="Each content item must be a dictionary"
)
item_type = content_item.get("type")
if not isinstance(item_type, str):
raise OCIError(
status_code=400,
message="Each content item must have a string `type` field",
)
if item_type not in ["text", "image_url"]:
raise OCIError(
status_code=400,
message=f"Content type `{item_type}` is not supported by OCI",
)
if item_type == "text":
text = content_item.get("text")
if not isinstance(text, str):
raise OCIError(
status_code=400,
message="Content item of type `text` must have a string `text` field",
)
new_content.append(OCITextContentPart(text=text))
elif item_type == "image_url":
image_url = content_item.get("image_url")
if isinstance(image_url, dict):
image_url = image_url.get("url")
if not isinstance(image_url, str):
raise OCIError(
status_code=400,
message="Prop `image_url` must be a string or an object with a `url` property",
)
new_content.append(OCIImageContentPart(imageUrl=OCIImageUrl(url=image_url)))
return OCIMessage(
role=open_ai_to_generic_oci_role_map[role],
content=new_content,
toolCalls=None,
toolCallId=None,
)
def adapt_messages_to_generic_oci_standard_tool_call(
role: str, tool_calls: list
) -> OCIMessage:
"""Convert an assistant tool-call message to OCI format."""
tool_calls_formatted = []
for tool_call in tool_calls:
if not isinstance(tool_call, dict):
raise OCIError(
status_code=400, message="Each tool call must be a dictionary"
)
if tool_call.get("type") != "function":
raise OCIError(
status_code=400, message="OCI only supports function tool calls"
)
tool_call_id = tool_call.get("id")
if not isinstance(tool_call_id, str):
raise OCIError(status_code=400, message="Tool call `id` must be a string")
tool_function = tool_call.get("function")
if not isinstance(tool_function, dict):
raise OCIError(
status_code=400, message="Tool call `function` must be a dictionary"
)
function_name = tool_function.get("name")
if not isinstance(function_name, str):
raise OCIError(
status_code=400, message="Tool call `function.name` must be a string"
)
arguments = tool_call["function"].get("arguments", "{}")
if not isinstance(arguments, str):
raise OCIError(
status_code=400,
message="Tool call `function.arguments` must be a JSON string",
)
tool_calls_formatted.append(
OCIToolCall(
id=tool_call_id,
type="FUNCTION",
name=function_name,
arguments=arguments,
)
)
return OCIMessage(
role=open_ai_to_generic_oci_role_map[role],
content=None,
toolCalls=tool_calls_formatted,
toolCallId=None,
)
def adapt_messages_to_generic_oci_standard_tool_response(
role: str, tool_call_id: str, content: str
) -> OCIMessage:
"""Convert a tool-result message to OCI format."""
return OCIMessage(
role=open_ai_to_generic_oci_role_map[role],
content=[OCITextContentPart(text=content)],
toolCalls=None,
toolCallId=tool_call_id,
)
def adapt_messages_to_generic_oci_standard(
messages: List[AllMessageValues],
) -> List[OCIMessage]:
"""Convert an OpenAI-format message array to OCI GENERIC format."""
new_messages = []
for message in messages:
role = message["role"]
content = message.get("content")
tool_calls = message.get("tool_calls")
tool_call_id = message.get("tool_call_id")
if role == "assistant" and tool_calls is not None:
if not isinstance(tool_calls, list):
raise OCIError(
status_code=400, message="Message `tool_calls` must be a list"
)
new_messages.append(
adapt_messages_to_generic_oci_standard_tool_call(role, tool_calls)
)
elif role in ["system", "user", "assistant"] and content is not None:
if not isinstance(content, (str, list)):
raise OCIError(
status_code=400,
message="Message `content` must be a string or list of content parts",
)
new_messages.append(
adapt_messages_to_generic_oci_standard_content_message(role, content)
)
elif role == "tool":
if not isinstance(tool_call_id, str):
raise OCIError(
status_code=400,
message="Tool result message must have a string `tool_call_id`",
)
if not isinstance(content, str):
raise OCIError(
status_code=400,
message="Tool result message `content` must be a string",
)
new_messages.append(
adapt_messages_to_generic_oci_standard_tool_response(
role, tool_call_id, content
)
)
return new_messages
# ---------------------------------------------------------------------------
# Tool definition adaptation
# ---------------------------------------------------------------------------
def adapt_tool_definition_to_oci_standard(
tools: List[Dict], vendor: OCIVendors
) -> List[OCIToolDefinition]:
"""Convert OpenAI-format tool definitions to OCI GENERIC format.
Resolves ``$ref``/``$defs`` and ``anyOf`` that the OCI endpoint rejects.
"""
new_tools = []
for tool in tools:
if tool["type"] != "function":
raise OCIError(status_code=400, message="OCI only supports function tools")
tool_function = tool.get("function")
if not isinstance(tool_function, dict):
raise OCIError(
status_code=400, message="Tool `function` must be a dictionary"
)
raw_params = tool_function.get("parameters", {})
resolved_params = sanitize_oci_schema(
resolve_oci_schema_anyof(resolve_oci_schema_refs(raw_params))
)
new_tools.append(
OCIToolDefinition(
type="FUNCTION",
name=tool_function.get("name"),
description=tool_function.get("description", ""),
parameters=resolved_params,
)
)
return new_tools
def _normalize_oci_finish_reason(raw: Optional[str]) -> Optional[str]:
"""Map an OCI-specific finish reason to its OpenAI-standard equivalent.
OCI emits ``COMPLETE`` / ``MAX_TOKENS`` / ``TOOL_CALL(S)`` plus a long tail
of error/cancel reasons (``ERROR``, ``ERROR_TOXIC``, ``ERROR_LIMIT``,
``USER_CANCEL``, ``CONTENT_FILTERED``, ``CANCELLED``, ...). The OpenAI
spec only defines ``stop`` / ``length`` / ``tool_calls`` / ... — anything
else is collapsed to ``"stop"`` so downstream consumers switching on
``finish_reason`` keep working. A ``None`` input passes through unchanged.
"""
if raw is None:
return None
if raw == "COMPLETE":
return "stop"
if raw == "MAX_TOKENS":
return "length"
if raw in ("TOOL_CALL", "TOOL_CALLS"):
return "tool_calls"
return "stop"
def _synthesize_oci_tool_call_id(position: int, name: str, arguments: str) -> str:
"""Deterministic synthetic tool-call id derived from chunk content.
Used as a fallback when OCI omits ``id`` (always the case for the OCI
Cohere protocol, occasionally the case for OCI GENERIC streaming chunks).
A random ``uuid4`` per chunk would cause downstream stream-merging
consumers — which key off the tool-call ``id`` — to treat re-emissions of
the same logical call (e.g. terminal consolidation chunks, retries) as
distinct calls. A content-derived digest stays stable across identical
re-emissions while differing across truly distinct calls.
"""
digest = hashlib.sha256(
f"{position}|{name}|{arguments}".encode("utf-8"),
usedforsecurity=False,
).hexdigest()[:24]
return f"call_{digest}"
def adapt_tools_to_openai_standard(
tools: List[OCIToolCall],
) -> List[ChatCompletionMessageToolCall]:
"""Convert OCI tool-call objects in a response to the OpenAI format."""
return [
ChatCompletionMessageToolCall(
id=tool.id or _synthesize_oci_tool_call_id(i, tool.name, tool.arguments),
type="function",
function={"name": tool.name, "arguments": tool.arguments},
)
for i, tool in enumerate(tools)
]
# ---------------------------------------------------------------------------
# Response parsing
# ---------------------------------------------------------------------------
def handle_generic_response(
json_data: dict,
model: str,
model_response: ModelResponse,
raw_response: httpx.Response,
) -> ModelResponse:
"""Parse a non-streaming GENERIC OCI response into a LiteLLM ModelResponse."""
try:
completion_response = OCICompletionResponse(**json_data)
except (TypeError, ValidationError) as e:
raise OCIError(
message=f"Response cannot be casted to OCICompletionResponse: {str(e)}",
status_code=raw_response.status_code,
)
iso_str = completion_response.chatResponse.timeCreated
dt = datetime.datetime.fromisoformat(iso_str.replace("Z", "+00:00"))
model_response.created = int(dt.timestamp())
model_response.model = completion_response.modelId
if not completion_response.chatResponse.choices:
raise OCIError(
message="OCI response contained no choices",
status_code=raw_response.status_code,
)
response_choice = completion_response.chatResponse.choices[0]
message = model_response.choices[0].message # type: ignore
response_message = response_choice.message
if response_message is not None:
if response_message.content:
# Concatenate all text parts — matches the streaming handler, which
# iterates the full content array. Skips non-text parts (e.g. image
# parts) so a leading non-text part doesn't suppress trailing text.
text: Optional[str] = None
for item in response_message.content:
if isinstance(item, OCITextContentPart):
text = (text or "") + item.text
if text is not None:
message.content = text
if response_message.toolCalls:
message.tool_calls = adapt_tools_to_openai_standard(
response_message.toolCalls
)
model_response.choices[0].finish_reason = _normalize_oci_finish_reason( # type: ignore[union-attr,assignment]
response_choice.finishReason
)
oci_usage = completion_response.chatResponse.usage
reasoning_tokens: Optional[int] = None
if (
oci_usage.completionTokensDetails
and oci_usage.completionTokensDetails.reasoningTokens is not None
):
reasoning_tokens = oci_usage.completionTokensDetails.reasoningTokens
model_response.usage = Usage( # type: ignore[attr-defined]
prompt_tokens=oci_usage.promptTokens,
completion_tokens=oci_usage.completionTokens or 0,
total_tokens=oci_usage.totalTokens,
reasoning_tokens=reasoning_tokens,
)
return model_response
def handle_generic_stream_chunk(dict_chunk: dict) -> ModelResponseStream:
"""Parse a single GENERIC SSE chunk into a LiteLLM ModelResponseStream."""
# OCI streams tool calls progressively — early chunks may omit required fields.
if dict_chunk.get("message") and dict_chunk["message"].get("toolCalls"):
for tool_call in dict_chunk["message"]["toolCalls"]:
tool_call.setdefault("arguments", "")
tool_call.setdefault("id", "")
tool_call.setdefault("name", "")
try:
typed_chunk = OCIStreamChunk(**dict_chunk)
except (TypeError, ValidationError) as e:
raise OCIError(
status_code=500,
message=f"Chunk cannot be parsed as OCIStreamChunk: {str(e)}",
)
if typed_chunk.index is None:
typed_chunk.index = 0
# Emit ``content=None`` rather than ``content=""`` on chunks with no text
# parts (e.g. tool-call-only or keep-alive chunks) so downstream
# stream-mergers that distinguish "no text in this delta" from "an
# explicitly empty text delta" behave correctly.
text: Optional[str] = None
if typed_chunk.message and typed_chunk.message.content:
for item in typed_chunk.message.content:
if isinstance(item, OCITextContentPart):
text = (text or "") + item.text
elif isinstance(item, OCIImageContentPart):
raise OCIError(
status_code=500,
message="OCI returned image content in a streaming response — not supported",
)
else:
raise OCIError(
status_code=500,
message=f"Unsupported content type in OCI streaming response: {item.type}",
)
# Build plain tool-call dicts inline (matching the shape produced by
# ``handle_cohere_stream_chunk``) rather than calling
# ``adapt_tools_to_openai_standard`` and ``model_dump``-ing the typed
# objects. Both code paths feed ``Delta.tool_calls``, so emitting the
# same minimal ``{"id", "type", "function": {"name", "arguments"}}``
# shape keeps downstream stream-mergers behaving identically across
# GENERIC and Cohere chunks.
tool_calls: Optional[List[Dict[str, Any]]] = None
if typed_chunk.message and typed_chunk.message.toolCalls:
tool_calls = [
{
"id": tc.id or _synthesize_oci_tool_call_id(i, tc.name, tc.arguments),
"type": "function",
"function": {
"name": tc.name,
"arguments": tc.arguments,
},
}
for i, tc in enumerate(typed_chunk.message.toolCalls)
]
finish_reason: Optional[str] = _normalize_oci_finish_reason(
typed_chunk.finishReason
)
return ModelResponseStream(
choices=[
StreamingChoices(
index=typed_chunk.index,
delta=Delta(
content=text,
tool_calls=tool_calls,
provider_specific_fields=None,
thinking_blocks=None,
reasoning_content=None,
),
finish_reason=finish_reason,
)
]
)

File diff suppressed because it is too large Load diff

View file

@ -1,9 +1,42 @@
from typing import Optional
import base64
import hashlib
import json
import os
import re
from dataclasses import dataclass
from email.utils import formatdate
from typing import Any, Dict, Optional, Protocol, Tuple
from urllib.parse import urlparse
import httpx
from litellm.llms.base_llm.chat.transformation import BaseLLMException
try:
from cryptography.hazmat.primitives import hashes, serialization
from cryptography.hazmat.primitives.asymmetric import padding, rsa
_CRYPTOGRAPHY_AVAILABLE = True
except ImportError:
_CRYPTOGRAPHY_AVAILABLE = False
try:
from litellm._version import version as _litellm_version
except ImportError:
_litellm_version = "0.0.0"
# OCI GenAI REST API version — stable since service launch, unlikely to change
OCI_API_VERSION = "20231130"
def _require_cryptography() -> None:
if not _CRYPTOGRAPHY_AVAILABLE:
raise ImportError(
"cryptography package is required for OCI authentication. "
"Please install it with: pip install cryptography"
)
class OCIError(BaseLLMException):
def __init__(
@ -17,3 +50,520 @@ class OCIError(BaseLLMException):
message=message,
headers=headers,
)
# ---------------------------------------------------------------------------
# OCI signing protocol and helpers
# ---------------------------------------------------------------------------
class OCISignerProtocol(Protocol):
"""
Protocol for OCI request signers (e.g., oci.signer.Signer).
Compatible with the OCI Python SDK's Signer class.
See: https://docs.oracle.com/en-us/iaas/tools/python/latest/api/signing.html
"""
def do_request_sign(
self, request: Any, *, enforce_content_headers: bool = False
) -> None:
pass
@dataclass
class OCIRequestWrapper:
"""
Wrapper for HTTP requests compatible with OCI signer interface.
Wraps request data in the format expected by OCI SDK signers, which require
objects with method, url, headers, body, and path_url attributes.
"""
method: str
url: str
headers: dict
body: bytes
@property
def path_url(self) -> str:
"""Returns the path + query string for OCI signing."""
parsed = urlparse(self.url)
return parsed.path + ("?" + parsed.query if parsed.query else "")
def sha256_base64(data: bytes) -> str:
# SHA-256 is used here to compute the x-content-sha256 header required by the
# OCI HTTP signing specification (RSA-SHA256 request signing), not for password
# or secret hashing. This is the correct and mandated algorithm for this purpose.
# See: https://docs.oracle.com/en-us/iaas/Content/API/Concepts/signingrequests.htm
#
# ``usedforsecurity=False`` declares non-security intent to static analyzers
# (CodeQL ``py/weak-sensitive-data-hashing``) — without it the request body
# gets flagged as "password-like data" via taint tracking.
digest = hashlib.sha256(data, usedforsecurity=False).digest() # noqa: S324
return base64.b64encode(digest).decode()
def build_signature_string(
method: str, path: str, headers: dict, signed_headers: list
) -> str:
lines = []
for header in signed_headers:
if header == "(request-target)":
value = f"{method.lower()} {path}"
else:
value = headers[header]
lines.append(f"{header}: {value}")
return "\n".join(lines)
def load_private_key_from_str(key_str: str) -> Any:
_require_cryptography()
key = serialization.load_pem_private_key( # type: ignore[union-attr]
key_str.encode("utf-8"),
password=None,
)
if not isinstance(key, rsa.RSAPrivateKey): # type: ignore[union-attr]
raise TypeError(
"The provided private key is not an RSA key, which is required for OCI signing."
)
return key
def load_private_key_from_file(file_path: str) -> Any:
"""Loads a private key from a file path."""
try:
with open(file_path, "r", encoding="utf-8") as f:
key_str = f.read().strip()
except FileNotFoundError:
raise FileNotFoundError(f"Private key file not found: {file_path}")
except OSError as e:
raise OSError(f"Failed to read private key file '{file_path}': {e}") from e
if not key_str:
raise ValueError(f"Private key file is empty: {file_path}")
return load_private_key_from_str(key_str)
# ---------------------------------------------------------------------------
# Env-var credential resolution
# ---------------------------------------------------------------------------
_OCI_REGION_ENV = "OCI_REGION"
_OCI_USER_ENV = "OCI_USER"
_OCI_FINGERPRINT_ENV = "OCI_FINGERPRINT"
_OCI_TENANCY_ENV = "OCI_TENANCY"
_OCI_KEY_FILE_ENV = "OCI_KEY_FILE"
_OCI_KEY_ENV = "OCI_KEY"
_OCI_COMPARTMENT_ID_ENV = "OCI_COMPARTMENT_ID"
def resolve_oci_credentials(optional_params: dict) -> dict:
"""
Merge OCI credentials from optional_params (explicit, always wins) and
environment variables (fallback).
Returns a dict with resolved values for:
oci_region, oci_user, oci_fingerprint, oci_tenancy,
oci_key, oci_key_file, oci_compartment_id
"""
return {
"oci_region": optional_params.get("oci_region")
or os.environ.get(_OCI_REGION_ENV)
or "us-ashburn-1",
"oci_user": optional_params.get("oci_user") or os.environ.get(_OCI_USER_ENV),
"oci_fingerprint": optional_params.get("oci_fingerprint")
or os.environ.get(_OCI_FINGERPRINT_ENV),
"oci_tenancy": optional_params.get("oci_tenancy")
or os.environ.get(_OCI_TENANCY_ENV),
"oci_key": optional_params.get("oci_key") or os.environ.get(_OCI_KEY_ENV),
"oci_key_file": optional_params.get("oci_key_file")
or os.environ.get(_OCI_KEY_FILE_ENV),
"oci_compartment_id": optional_params.get("oci_compartment_id")
or os.environ.get(_OCI_COMPARTMENT_ID_ENV),
}
_OCI_REGION_RE = re.compile(r"^[a-z][a-z0-9-]{0,30}[a-z0-9]$")
_OCI_ACTION_PATH_RE = re.compile(rf"/{OCI_API_VERSION}/actions/[^/?#]+/?$")
def get_oci_base_url(optional_params: dict, api_base: Optional[str] = None) -> str:
"""Return the OCI inference base URL, respecting any explicit api_base override.
If ``api_base`` already ends with a fully-formed OCI action path
(``/{OCI_API_VERSION}/actions/<name>``), that suffix is stripped so callers
can append their own action path without producing a doubled URL.
"""
if api_base:
return _OCI_ACTION_PATH_RE.sub("", api_base).rstrip("/")
creds = resolve_oci_credentials(optional_params)
region = creds["oci_region"]
if not isinstance(region, str) or not _OCI_REGION_RE.match(region):
raise OCIError(
status_code=400,
message=(
f"Invalid OCI region {region!r}: must match "
"^[a-z][a-z0-9-]{0,30}[a-z0-9]$ (e.g. 'us-ashburn-1')."
),
)
return f"https://inference.generativeai.{region}.oci.oraclecloud.com"
# ---------------------------------------------------------------------------
# Signing implementations (shared by chat, embed, and rerank configs)
# ---------------------------------------------------------------------------
def sign_with_oci_signer(
headers: dict,
optional_params: dict,
request_data: dict,
api_base: str,
) -> Tuple[dict, bytes]:
"""Sign a request using an OCI SDK Signer object passed in optional_params."""
oci_signer = optional_params.get("oci_signer")
body = json.dumps(request_data).encode("utf-8")
method = str(optional_params.get("method", "POST")).upper()
if method not in {"POST", "GET", "PUT", "DELETE", "PATCH"}:
raise ValueError(f"Unsupported HTTP method: {method}")
prepared_headers = {**headers}
prepared_headers.setdefault("content-type", "application/json")
prepared_headers.setdefault("content-length", str(len(body)))
request_wrapper = OCIRequestWrapper(
method=method, url=api_base, headers=prepared_headers, body=body
)
if oci_signer is None:
raise ValueError("oci_signer cannot be None when calling sign_with_oci_signer")
try:
oci_signer.do_request_sign(request_wrapper, enforce_content_headers=True)
except Exception as e:
raise OCIError(
status_code=500,
message=(
f"Failed to sign request with provided oci_signer: {str(e)}. "
"The signer must implement the OCI SDK Signer interface with a "
"do_request_sign(request, enforce_content_headers=True) method. "
"See: https://docs.oracle.com/en-us/iaas/tools/python/latest/api/signing.html"
),
) from e
headers.update(request_wrapper.headers)
return headers, body
def sign_with_manual_credentials(
headers: dict,
optional_params: dict,
request_data: dict,
api_base: str,
) -> Tuple[dict, bytes]:
"""Sign a request using manually provided OCI credentials (user/fingerprint/tenancy/key)."""
creds = resolve_oci_credentials(optional_params)
oci_user = creds["oci_user"]
oci_fingerprint = creds["oci_fingerprint"]
oci_tenancy = creds["oci_tenancy"]
oci_key = creds["oci_key"]
oci_key_file = creds["oci_key_file"]
if (
not oci_user
or not oci_fingerprint
or not oci_tenancy
or not (oci_key or oci_key_file)
):
raise OCIError(
status_code=401,
message=(
"Missing required OCI credentials: oci_user, oci_fingerprint, oci_tenancy, "
"and at least one of oci_key or oci_key_file. "
"These can also be supplied via environment variables: "
f"{_OCI_USER_ENV}, {_OCI_FINGERPRINT_ENV}, {_OCI_TENANCY_ENV}, {_OCI_KEY_ENV} (or {_OCI_KEY_FILE_ENV}). "
"Alternatively, provide an oci_signer object from the OCI SDK."
),
)
method = str(optional_params.get("method", "POST")).upper()
body = json.dumps(request_data).encode("utf-8")
parsed = urlparse(api_base)
path = parsed.path or "/"
host = parsed.netloc
date = formatdate(usegmt=True)
content_type = headers.get("content-type", "application/json")
content_length = str(len(body))
x_content_sha256 = sha256_base64(body)
headers_to_sign: Dict[str, str] = {
"date": date,
"host": host,
"content-type": content_type,
"content-length": content_length,
"x-content-sha256": x_content_sha256,
}
signed_header_names = [
"date",
"(request-target)",
"host",
"content-length",
"content-type",
"x-content-sha256",
]
signing_string = build_signature_string(
method, path, headers_to_sign, signed_header_names
)
_require_cryptography()
# Resolve the private key — prefer inline PEM content over file path
oci_key_content: Optional[str] = None
if oci_key:
if not isinstance(oci_key, str):
raise OCIError(
status_code=400,
message=(
f"oci_key must be a string containing the PEM private key content. "
f"Got type: {type(oci_key).__name__}"
),
)
oci_key_content = oci_key.replace("\\n", "\n").replace("\r\n", "\n")
private_key = (
load_private_key_from_str(oci_key_content)
if oci_key_content
else load_private_key_from_file(oci_key_file) if oci_key_file else None
)
if private_key is None:
raise OCIError(
status_code=400,
message="Private key is required for OCI authentication. Provide either oci_key or oci_key_file.",
)
signature = private_key.sign(
signing_string.encode("utf-8"),
padding.PKCS1v15(), # type: ignore[union-attr]
hashes.SHA256(), # type: ignore[union-attr]
)
signature_b64 = base64.b64encode(signature).decode()
key_id = f"{oci_tenancy}/{oci_user}/{oci_fingerprint}"
authorization = (
'Signature version="1",'
f'keyId="{key_id}",'
'algorithm="rsa-sha256",'
f'headers="{" ".join(signed_header_names)}",'
f'signature="{signature_b64}"'
)
headers.update(
{
"authorization": authorization,
"date": date,
"host": host,
"content-type": content_type,
"content-length": content_length,
"x-content-sha256": x_content_sha256,
}
)
return headers, body
def sign_oci_request(
headers: dict,
optional_params: dict,
request_data: dict,
api_base: str,
api_key: Optional[str] = None,
model: Optional[str] = None,
stream: Optional[bool] = None,
fake_stream: Optional[bool] = None,
) -> Tuple[dict, bytes]:
"""
Route to the appropriate OCI signing method based on what credentials are present.
If ``oci_signer`` is in optional_params, use the OCI SDK signer object.
Otherwise use manual RSA-SHA256 signing with explicit credentials (which can
also be supplied via OCI_* environment variables).
Returns:
Tuple of (signed_headers, signed_body_bytes)
"""
if optional_params.get("oci_signer") is not None:
return sign_with_oci_signer(headers, optional_params, request_data, api_base)
return sign_with_manual_credentials(
headers, optional_params, request_data, api_base
)
def validate_oci_environment(
headers: dict,
optional_params: dict,
api_key: Optional[str] = None,
) -> dict:
"""
Populate common OCI request headers (content-type, user-agent).
Full credential validation is deferred to signing time so that credentials
supplied via environment variables are resolved at call time rather than
at construction time.
"""
headers.setdefault("content-type", "application/json")
headers.setdefault("user-agent", f"litellm/{_litellm_version}")
return headers
# ---------------------------------------------------------------------------
# JSON schema utilities for OCI tool definitions
#
# OCI Generative AI does not support JSON Schema extensions ($ref, $defs,
# anyOf). Pydantic v2 emits all three for models with Optional fields or
# nested schemas. The helpers below are ported from the official
# langchain-oracle reference implementation so that tool schemas are always
# valid before they reach the OCI endpoint.
# ---------------------------------------------------------------------------
# Mapping from JSON Schema type names to Python type names, as expected by
# the OCI Cohere API's CohereParameterDefinition.type field.
OCI_JSON_TO_PYTHON_TYPES: Dict[str, str] = {
"string": "str",
"number": "float",
"boolean": "bool",
"integer": "int",
"array": "List",
"object": "Dict",
"any": "any",
}
def resolve_oci_schema_refs(schema: Dict[str, Any]) -> Dict[str, Any]:
"""Inline all ``$ref``/``$defs`` references — OCI does not support JSON Schema ``$ref``."""
defs = schema.get("$defs", {})
resolving_stack: set = set()
def _resolve(obj: Any) -> Any:
if isinstance(obj, dict):
if "$ref" in obj:
ref = obj["$ref"]
if ref.startswith("#/$defs/"):
key = ref.split("/")[-1]
if key in resolving_stack:
return {"type": "object"} # break cycles
resolving_stack.add(key)
try:
return _resolve(defs.get(key, obj))
finally:
resolving_stack.discard(key)
return obj # external $ref — leave unchanged
return {k: _resolve(v) for k, v in obj.items()}
if isinstance(obj, list):
return [_resolve(item) for item in obj]
return obj
resolved = _resolve(schema)
if isinstance(resolved, dict):
resolved.pop("$defs", None)
return resolved
def resolve_oci_schema_anyof(obj: Any) -> Any:
"""Resolve Pydantic v2 ``Optional[T]`` → ``anyOf`` patterns.
Pydantic v2 emits ``{"anyOf": [{"type": "T"}, {"type": "null"}]}`` for
``Optional[T]``. OCI models don't understand ``anyOf``, so we pick the
first non-null branch and merge top-level metadata into it.
"""
if isinstance(obj, dict):
if "anyOf" in obj and "type" not in obj:
non_null = [
t
for t in obj["anyOf"]
if not (isinstance(t, dict) and t.get("type") == "null")
]
if non_null:
resolved = {**obj, **non_null[0]}
resolved.pop("anyOf", None)
return resolve_oci_schema_anyof(resolved)
return {k: resolve_oci_schema_anyof(v) for k, v in obj.items()}
if isinstance(obj, list):
return [resolve_oci_schema_anyof(item) for item in obj]
return obj
def sanitize_oci_schema(schema: Any) -> Any:
"""Recursively remove OCI-incompatible fields from a JSON schema.
Strips ``title`` keys, removes ``None``-valued ``default`` entries,
normalises ``type: [T, "null"]`` list types, and ensures arrays carry an
``items`` definition.
"""
if isinstance(schema, list):
return [sanitize_oci_schema(item) for item in schema]
if not isinstance(schema, dict):
return schema
sanitized: Dict[str, Any] = {}
for key, value in schema.items():
if key == "title":
continue
if key == "default" and value is None:
continue
if key == "type":
if value == "any":
sanitized[key] = "object"
continue
if isinstance(value, list):
non_null = [t for t in value if t != "null"]
sanitized[key] = non_null[0] if non_null else "string"
continue
sanitized[key] = sanitize_oci_schema(value)
if sanitized.get("type") == "array" and "items" not in sanitized:
sanitized["items"] = {"type": "object"}
required = sanitized.get("required")
properties = sanitized.get("properties")
if "required" in sanitized:
if isinstance(required, list) and isinstance(properties, dict):
sanitized["required"] = [
f for f in required if isinstance(f, str) and f in properties
]
elif not isinstance(required, list):
sanitized["required"] = []
return sanitized
def enrich_cohere_param_description(
description: str, param_schema: Dict[str, Any]
) -> str:
"""Embed schema constraints into a Cohere parameter description.
``CohereParameterDefinition`` only has ``type``, ``description``, and
``isRequired``. Rich constraints (``enum``, ``format``, ``minimum``,
``maximum``, ``pattern``) are appended to the description string so the
model can still see and respect them.
"""
parts = [description] if description else []
if "enum" in param_schema:
parts.append(f"Allowed values: {param_schema['enum']}")
if "format" in param_schema:
parts.append(f"Format: {param_schema['format']}")
if "minimum" in param_schema or "maximum" in param_schema:
range_parts = []
if "minimum" in param_schema:
range_parts.append(f"min={param_schema['minimum']}")
if "maximum" in param_schema:
range_parts.append(f"max={param_schema['maximum']}")
parts.append(f"Range: {', '.join(range_parts)}")
if "pattern" in param_schema:
parts.append(f"Pattern: {param_schema['pattern']}")
return ". ".join(parts) if parts else ""

View file

@ -1,8 +1,14 @@
"""
OCI Generative AI Embedding Configuration
OCI Generative AI — Embedding transformation.
Supports embedding models available on Oracle Cloud Infrastructure Generative AI service.
Uses the same authentication mechanisms as OCI chat (manual signing or OCI SDK Signer).
Endpoint: POST /20231130/actions/embedText
Supported models: cohere.embed-english-v3.0, cohere.embed-multilingual-v3.0,
cohere.embed-v4.0, and all other Cohere embed variants available on OCI
(including dedicated endpoints).
Authentication follows the same RSA-SHA256 / OCI SDK signer pattern as chat.
The base handler (base_llm_http_handler.embedding) calls sign_request after
building the body, so signing happens automatically.
Supported models:
- cohere.embed-english-v3.0
@ -10,25 +16,45 @@ Supported models:
- cohere.embed-multilingual-v3.0
- cohere.embed-multilingual-light-v3.0
- cohere.embed-english-image-v3.0
- cohere.embed-english-light-image-v3.0
- cohere.embed-multilingual-light-image-v3.0
- cohere.embed-multilingual-image-v3.0
- cohere.embed-v4.0
Reference: https://docs.oracle.com/en-us/iaas/api/#/en/generative-ai-inference/latest/EmbedTextResult/EmbedText
"""
from typing import Any, Dict, List, Optional, Union
from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union
import httpx
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
import litellm
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
from litellm.llms.oci.chat.transformation import OCIChatConfig
from litellm.llms.oci.common_utils import OCIError
from litellm.llms.oci.common_utils import (
OCI_API_VERSION,
OCIError,
get_oci_base_url,
resolve_oci_credentials,
sign_oci_request,
validate_oci_environment,
)
from litellm.types.llms.oci import (
OCIEmbedRequest,
OCIEmbedResponse,
OCIServingMode,
)
from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues
from litellm.types.utils import EmbeddingResponse, Usage
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
# OCI sends up to 96 texts per embedText request (Cohere limit).
OCI_EMBED_BATCH_LIMIT = 96
# Input type mapping from OpenAI conventions to OCI/Cohere conventions
_INPUT_TYPE_MAP = {
"search_document": "SEARCH_DOCUMENT",
@ -38,65 +64,43 @@ _INPUT_TYPE_MAP = {
}
class OCIEmbeddingConfig(BaseEmbeddingConfig):
class OCIEmbedConfig(BaseEmbeddingConfig):
"""
Configuration for OCI Generative AI Embedding API.
Transformation config for OCI Generative AI embeddings.
The OCI embedding endpoint uses the Cohere embed models hosted on OCI.
Authentication is handled via OCI request signing (manual credentials or OCI SDK Signer).
Supports both text and (on cohere.embed-v4.0) multimodal inputs.
Usage:
```python
import litellm
Authentication — same two modes as chat:
- **OCI SDK signer**: pass ``oci_signer`` in optional_params.
- **Manual RSA-SHA256**: pass ``oci_user``, ``oci_fingerprint``, ``oci_tenancy``,
and ``oci_key`` or ``oci_key_file``, or set the corresponding ``OCI_*`` env vars.
response = litellm.embedding(
model="oci/cohere.embed-english-v3.0",
input=["Hello world", "Goodbye world"],
oci_compartment_id="ocid1.compartment.oc1..xxx",
oci_region="us-ashburn-1",
oci_user="ocid1.user.oc1..xxx",
oci_fingerprint="xx:xx:xx:xx",
oci_tenancy="ocid1.tenancy.oc1..xxx",
oci_key_file="~/.oci/key.pem",
)
```
Required call-time params (via optional_params or env vars):
- ``oci_compartment_id`` / ``OCI_COMPARTMENT_ID``
- ``oci_region`` / ``OCI_REGION`` (default: ``us-ashburn-1``)
Optional call-time params:
- ``oci_serving_mode``: ``"ON_DEMAND"`` (default) or ``"DEDICATED"``
- ``oci_endpoint_id``: endpoint OCID for dedicated serving mode
- ``input_type``: ``SEARCH_DOCUMENT``, ``SEARCH_QUERY``, ``CLASSIFICATION``, ``CLUSTERING``
- ``truncate``: ``NONE``, ``START``, or ``END`` (default ``END``)
- ``dimensions``: output embedding dimensions (cohere.embed-v4.0+)
"""
def __init__(self) -> None:
# We reuse OCIChatConfig for signing logic
self._chat_config = OCIChatConfig()
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:
if api_base:
return api_base
oci_region = optional_params.get("oci_region", "us-ashburn-1")
return f"https://inference.generativeai.{oci_region}.oci.oraclecloud.com/20231130/actions/embedText"
def get_supported_openai_params(self, model: str) -> list:
return [
"dimensions",
]
def get_supported_openai_params(self, model: str) -> List[str]:
return ["dimensions"]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
drop_params: bool = False,
) -> dict:
# Note: OCI Cohere embed does not support custom dimensions natively,
# but we pass it through in case future models support it
if "dimensions" in non_default_params:
optional_params["dimensions"] = non_default_params["dimensions"]
for key, value in non_default_params.items():
if key == "dimensions":
# OCI API uses outputDimensions (cohere.embed-v4.0+)
optional_params["outputDimensions"] = value
return optional_params
def validate_environment(
@ -109,49 +113,42 @@ class OCIEmbeddingConfig(BaseEmbeddingConfig):
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> dict:
"""
Validate OCI credentials for embedding requests.
Supports both OCI SDK Signer and manual credential signing.
"""
oci_signer = optional_params.get("oci_signer")
oci_region = optional_params.get("oci_region", "us-ashburn-1")
api_base = (
api_base
or f"https://inference.generativeai.{oci_region}.oci.oraclecloud.com"
)
if oci_signer is None:
oci_user = optional_params.get("oci_user")
oci_fingerprint = optional_params.get("oci_fingerprint")
oci_tenancy = optional_params.get("oci_tenancy")
oci_key = optional_params.get("oci_key")
oci_key_file = optional_params.get("oci_key_file")
oci_compartment_id = optional_params.get("oci_compartment_id")
if (
not oci_user
or not oci_fingerprint
or not oci_tenancy
or not (oci_key or oci_key_file)
or not oci_compartment_id
):
raise Exception(
"Missing required parameters: oci_user, oci_fingerprint, oci_tenancy, oci_compartment_id "
"and at least one of oci_key or oci_key_file. "
"Alternatively, provide an oci_signer object from the OCI SDK."
if optional_params.get("oci_signer") is None:
creds = resolve_oci_credentials(optional_params)
missing = [
k
for k in (
"oci_user",
"oci_fingerprint",
"oci_tenancy",
"oci_compartment_id",
)
if not creds.get(k)
]
if missing or not (creds.get("oci_key") or creds.get("oci_key_file")):
raise OCIError(
status_code=401,
message=(
"Missing required parameters: oci_user, oci_fingerprint, oci_tenancy, "
"oci_compartment_id and at least one of oci_key or oci_key_file. "
"These can be supplied via optional_params or via OCI_USER, OCI_FINGERPRINT, "
"OCI_TENANCY, OCI_COMPARTMENT_ID, OCI_KEY_FILE environment variables. "
"Alternatively, provide an oci_signer object from the OCI SDK."
),
)
return validate_oci_environment(headers, optional_params, api_key)
from litellm.llms.custom_httpx.http_handler import version
headers.update(
{
"content-type": "application/json",
"user-agent": f"litellm/{version}",
}
)
return 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 = get_oci_base_url(optional_params, api_base or litellm.api_base)
return f"{base}/{OCI_API_VERSION}/actions/embedText"
def sign_request(
self,
@ -163,9 +160,8 @@ class OCIEmbeddingConfig(BaseEmbeddingConfig):
model: Optional[str] = None,
stream: Optional[bool] = None,
fake_stream: Optional[bool] = None,
):
"""Delegate to OCIChatConfig's signing logic."""
return self._chat_config.sign_request(
) -> Tuple[dict, bytes]:
return sign_oci_request(
headers=headers,
optional_params=optional_params,
request_data=request_data,
@ -182,91 +178,74 @@ class OCIEmbeddingConfig(BaseEmbeddingConfig):
input: AllEmbeddingInputValues,
optional_params: dict,
headers: dict,
api_base: Optional[str] = None,
) -> dict:
"""
Transform the embedding request to OCI format.
OCI embedText API expects:
{
"compartmentId": "...",
"servingMode": {"servingType": "ON_DEMAND", "modelId": "..."},
"inputs": ["text1", "text2"],
"truncate": "END",
"inputType": "SEARCH_DOCUMENT"
}
"""
oci_compartment_id = optional_params.get("oci_compartment_id")
if not oci_compartment_id:
raise Exception(
"kwarg `oci_compartment_id` is required for OCI embedding requests"
creds = resolve_oci_credentials(optional_params)
compartment_id = creds["oci_compartment_id"]
if not compartment_id:
raise OCIError(
status_code=400,
message=(
"oci_compartment_id is required for OCI embedding requests. "
"Pass it as optional_params or set the OCI_COMPARTMENT_ID env var."
),
)
# Build serving mode
oci_serving_mode = optional_params.get("oci_serving_mode", "ON_DEMAND")
if oci_serving_mode == "DEDICATED":
oci_endpoint_id = optional_params.get("oci_endpoint_id", model)
serving_mode = {
"servingType": "DEDICATED",
"endpointId": oci_endpoint_id,
}
else:
serving_mode = {
"servingType": "ON_DEMAND",
"modelId": model,
}
# Normalize input to list of strings
# Normalise input to a flat list of strings
if isinstance(input, str):
inputs = [input]
texts = [input]
elif isinstance(input, list):
inputs = []
texts = []
for item in input:
if isinstance(item, str):
inputs.append(item)
elif isinstance(item, list):
raise ValueError(
"OCI embedding does not support token-array inputs. "
"Please convert token lists to strings before calling embedding()."
if isinstance(item, list):
raise OCIError(
status_code=400,
message=(
"OCI embedText does not support token-array inputs. "
"Convert token lists to strings before calling embedding()."
),
)
else:
inputs.append(str(item))
texts.append(item if isinstance(item, str) else str(item))
else:
inputs = [str(input)]
texts = [str(input)]
# Build request data — OCI embedText API expects inputs, truncate,
# and inputType at the top level alongside compartmentId and servingMode
request_data: Dict[str, Any] = {
"compartmentId": oci_compartment_id,
"servingMode": serving_mode,
"inputs": inputs,
"truncate": optional_params.get("truncate", "END"),
}
if len(texts) > OCI_EMBED_BATCH_LIMIT:
raise OCIError(
status_code=400,
message=(
f"OCI embedText accepts at most {OCI_EMBED_BATCH_LIMIT} inputs per request "
f"(got {len(texts)}). Batch your requests."
),
)
# Map input_type if provided
serving_mode_type = optional_params.get("oci_serving_mode", "ON_DEMAND").upper()
if serving_mode_type not in {"ON_DEMAND", "DEDICATED"}:
raise OCIError(
status_code=400,
message="oci_serving_mode must be 'ON_DEMAND' or 'DEDICATED'.",
)
if serving_mode_type == "DEDICATED":
endpoint_id = optional_params.get("oci_endpoint_id", model)
serving_mode = OCIServingMode(
servingType="DEDICATED", endpointId=endpoint_id
)
else:
serving_mode = OCIServingMode(servingType="ON_DEMAND", modelId=model)
# Map input_type from OpenAI convention to OCI/Cohere convention
input_type = optional_params.get("input_type")
if input_type:
mapped_type = _INPUT_TYPE_MAP.get(input_type.lower(), input_type.upper())
request_data["inputType"] = mapped_type
input_type = _INPUT_TYPE_MAP.get(input_type.lower(), input_type.upper())
# Sign the request using the same URL the HTTP handler will POST to
signing_url = self.get_complete_url(
api_base=api_base,
api_key=None,
model=model,
optional_params=optional_params,
litellm_params={},
request = OCIEmbedRequest(
compartmentId=compartment_id,
servingMode=serving_mode,
inputs=texts,
inputType=input_type,
truncate=optional_params.get("truncate", "END"),
outputDimensions=optional_params.get("outputDimensions"),
)
signed_headers, body = self.sign_request(
headers=headers,
optional_params=optional_params,
request_data=request_data,
api_base=signing_url,
)
headers.update(signed_headers)
return request_data
return request.model_dump(exclude_none=True)
def transform_embedding_response(
self,
@ -274,63 +253,57 @@ class OCIEmbeddingConfig(BaseEmbeddingConfig):
raw_response: httpx.Response,
model_response: EmbeddingResponse,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str] = None,
request_data: dict = {},
optional_params: dict = {},
litellm_params: dict = {},
api_key: Optional[str],
request_data: dict,
optional_params: dict,
litellm_params: dict,
) -> EmbeddingResponse:
"""
Transform OCI embedding response to standard EmbeddingResponse format.
OCI response format:
{
"embeddings": [[0.1, 0.2, ...], [0.3, 0.4, ...]],
"modelId": "cohere.embed-english-v3.0",
"modelVersion": "3.0",
"inputTextTokenCounts": [5, 4]
}
"""
if raw_response.status_code != 200:
raise OCIError(
message=raw_response.text,
status_code=raw_response.status_code,
message=raw_response.text,
)
try:
raw_response_json = raw_response.json()
except Exception:
json_response = raw_response.json()
except Exception as e:
raise OCIError(
message=raw_response.text,
status_code=raw_response.status_code,
message=f"Failed to parse OCI embed response as JSON: {e}",
)
embeddings = raw_response_json.get("embeddings", [])
model_id = raw_response_json.get("modelId", model)
# Build response data in OpenAI format
embedding_data = []
for idx, embedding in enumerate(embeddings):
embedding_data.append(
{
"object": "embedding",
"index": idx,
"embedding": embedding,
}
try:
parsed = OCIEmbedResponse(**json_response)
except Exception as e:
raise OCIError(
status_code=500,
message=f"OCI embed response does not match expected schema: {e}",
)
model_response.model = model_id
model_response.data = embedding_data
model_response.object = "list"
model_response.model = parsed.modelId
model_response.data = [
{
"object": "embedding",
"index": i,
"embedding": embedding,
}
for i, embedding in enumerate(parsed.embeddings)
]
# Calculate token usage
input_token_counts = raw_response_json.get("inputTextTokenCounts", [])
total_tokens = sum(input_token_counts) if input_token_counts else 0
usage = Usage(
prompt_tokens=total_tokens,
total_tokens=total_tokens,
)
model_response.usage = usage
if parsed.inputTextTokenCounts is not None:
# Actual OCI API returns per-input token counts — sum for total usage
total = sum(parsed.inputTextTokenCounts)
model_response.usage = Usage(prompt_tokens=total, total_tokens=total)
elif parsed.usage is not None:
# Some deployments may return a usage object directly
model_response.usage = Usage(
prompt_tokens=parsed.usage.promptTokens,
total_tokens=parsed.usage.totalTokens,
)
else:
# Neither field returned — default to zero so downstream consumers
# can always rely on usage being populated.
model_response.usage = Usage(prompt_tokens=0, total_tokens=0)
return model_response
@ -340,8 +313,8 @@ class OCIEmbeddingConfig(BaseEmbeddingConfig):
status_code: int,
headers: Union[dict, httpx.Headers],
) -> BaseLLMException:
return OCIError(
message=error_message,
status_code=status_code,
headers=headers if isinstance(headers, httpx.Headers) else None,
)
return OCIError(status_code=status_code, message=error_message)
# Alias for backwards compatibility with any code that imports OCIEmbeddingConfig
OCIEmbeddingConfig = OCIEmbedConfig

View file

@ -1,14 +1,14 @@
"""
Support for o1/o3 model family
Support for o1/o3 model family
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)
- Logprobs => 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)
"""
from typing import Any, Coroutine, List, Literal, Optional, Union, cast, overload

View file

@ -201,7 +201,7 @@ class BaseOpenAILLM:
@staticmethod
def get_openai_client_initialization_param_fields(
client_type: Literal["openai", "azure"]
client_type: Literal["openai", "azure"],
) -> Tuple[str, ...]:
"""Returns a tuple of fields that are used to initialize the OpenAI client"""
if client_type == "openai":

View file

@ -19,7 +19,10 @@ def cost_router(call_type: CallTypes) -> Literal["cost_per_token", "cost_per_sec
def cost_per_token(
model: str, usage: Usage, service_tier: Optional[str] = None
model: str,
usage: Usage,
service_tier: Optional[str] = None,
data_residency: Optional[str] = None,
) -> Tuple[float, float]:
"""
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
@ -27,6 +30,9 @@ def cost_per_token(
Input:
- model: str, the model name without provider prefix
- usage: LiteLLM Usage block, containing anthropic caching information
- data_residency: optional OpenAI data-residency region (e.g. "eu", "us"),
inferred from api_base. Applies the model's regional-processing
uplift multiplier when set.
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
@ -37,6 +43,7 @@ def cost_per_token(
usage=usage,
custom_llm_provider="openai",
service_tier=service_tier,
data_residency=data_residency,
)
# ### Non-cached text tokens
# non_cached_text_tokens = usage.prompt_tokens

View file

@ -0,0 +1,41 @@
"""
Helpers for resolving OpenAI data-residency (regional processing) from an
api_base URL.
OpenAI enforces hostname-per-region for projects with geography restrictions
enabled and rejects requests sent to the wrong host, so the api_base hostname
is the authoritative signal of which region a request was processed in.
"""
from typing import Dict, Optional
from urllib.parse import urlparse
# Mapping of OpenAI regional hostnames to the corresponding data-residency
# value used by the cost calculator. See
# https://developers.openai.com/api/docs/pricing for the regional-processing
# uplift these hostnames trigger.
_OPENAI_REGIONAL_HOSTS: Dict[str, str] = {
"eu.api.openai.com": "eu",
"us.api.openai.com": "us",
}
def infer_openai_data_residency(
custom_llm_provider: Optional[str], api_base: Optional[str]
) -> Optional[str]:
"""
Derive the OpenAI data-residency region from an api_base URL.
Returns ``"eu"`` for the EU regional host, ``"us"`` for the US regional
host, and ``None`` for the default global host, any non-OpenAI provider,
or any non-OpenAI URL.
"""
if custom_llm_provider != "openai" or not api_base:
return None
try:
host = urlparse(api_base).hostname
except (TypeError, ValueError):
return None
if not host:
return None
return _OPENAI_REGIONAL_HOSTS.get(host.lower())

View file

@ -126,9 +126,21 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Dict:
"""No transform applied since inputs are in OpenAI spec already"""
"""Strip Anthropic-only `cache_control` markers before sending to OpenAI.
OpenAI's Responses API rejects unknown fields on input content blocks
with HTTP 400 ("Unknown parameter: 'input[0].content[0].cache_control'").
Chat Completions strips these in
`remove_cache_control_flag_from_messages_and_tools`; mirror that here.
"""
input = self._validate_input_param(input)
tools = response_api_optional_request_params.get("tools")
input, tools = self.remove_cache_control_flag_from_input_and_tools(
model=model, input=input, tools=tools
)
if tools is not None:
response_api_optional_request_params["tools"] = tools
final_request_params = dict(
ResponsesAPIRequestParams(
model=model, input=input, **response_api_optional_request_params
@ -137,6 +149,38 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
return final_request_params
def remove_cache_control_flag_from_input_and_tools(
self,
model: str, # allows overrides to selectively run this
input: Union[str, ResponseInputParam],
tools: Optional[List[ALL_RESPONSES_API_TOOL_PARAMS]] = None,
) -> Tuple[
Union[str, ResponseInputParam],
Optional[List[ALL_RESPONSES_API_TOOL_PARAMS]],
]:
"""Sibling of `remove_cache_control_flag_from_messages_and_tools` on
the chat path. Strips Anthropic-only `cache_control` markers from
Responses API input content blocks and tools.
`filter_value_from_dict` mutates each dict in place, so the same
objects are returned.
"""
from litellm.litellm_core_utils.prompt_templates.common_utils import (
filter_value_from_dict,
)
if isinstance(input, list):
for item in input:
if isinstance(item, dict):
filter_value_from_dict(cast(dict, item), "cache_control")
if tools is not None:
for tool in tools:
if isinstance(tool, dict):
filter_value_from_dict(cast(dict, tool), "cache_control")
return input, tools
def _validate_input_param(
self, input: Union[str, ResponseInputParam]
) -> Union[str, ResponseInputParam]:
@ -604,6 +648,12 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
url = str(parsed_url.copy_with(path=compact_path))
input = self._validate_input_param(input)
tools = response_api_optional_request_params.get("tools")
input, tools = self.remove_cache_control_flag_from_input_and_tools(
model=model, input=input, tools=tools
)
if tools is not None:
response_api_optional_request_params["tools"] = tools
data = dict(
ResponsesAPIRequestParams(
model=model, input=input, **response_api_optional_request_params

View file

@ -49,7 +49,6 @@ from litellm.types.utils import (
)
from litellm.llms.openrouter.common_utils import OpenRouterException
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
else:

View file

@ -578,7 +578,7 @@ class SagemakerLLM(BaseAWSLLM):
logger_fn=None,
):
"""
Supports both Huggingface Jumpstart embeddings and Voyage models
Supports Hugging Face (TGI), Voyage, and Cohere embedding endpoints
"""
### BOTO3 INIT
import boto3

View file

@ -1,7 +1,7 @@
"""
Translate from OpenAI's `/v1/chat/completions` to Sagemaker's `/invoke`
In the Huggingface TGI format.
In the Huggingface TGI format.
"""
import json

View file

@ -0,0 +1,141 @@
"""
Translate from OpenAI's `/v1/embeddings` to Sagemaker's `/invoke`
In the native Cohere embed format for self-hosted Cohere endpoints
(AWS Marketplace / JumpStart). Cohere containers expect
`{"texts": [...], "input_type": "..."}` and reject the HuggingFace TGI shape
`{"inputs": [...]}` with `422 EmbedReqV2.inputs is of type string but should
be of type Object`.
Reference: https://docs.cohere.com/v2/reference/embed
"""
from typing import TYPE_CHECKING, Any, List, Optional, Union, cast
if TYPE_CHECKING:
from litellm.types.llms.openai import AllEmbeddingInputValues
from httpx._models import Headers, Response
import litellm
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
from litellm.llms.bedrock.embed.cohere_transformation import (
BedrockCohereEmbeddingConfig,
)
from litellm.llms.cohere.embed.v1_transformation import CohereEmbeddingConfig
from litellm.types.utils import EmbeddingResponse
from ..common_utils import SagemakerError
class SagemakerCohereEmbeddingConfig(BaseEmbeddingConfig):
"""
SageMaker invoke payload for self-hosted Cohere embed models.
"""
def __init__(self) -> None:
pass
def get_supported_openai_params(self, model: str) -> List[str]:
return ["encoding_format", "dimensions", "input_type"]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
optional_params = BedrockCohereEmbeddingConfig().map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
)
if "input_type" in non_default_params:
optional_params["input_type"] = non_default_params["input_type"]
return optional_params
def get_error_class(
self, error_message: str, status_code: int, headers: Union[dict, Headers]
) -> BaseLLMException:
return SagemakerError(
message=error_message, status_code=status_code, headers=headers
)
def transform_embedding_request(
self,
model: str,
input: "AllEmbeddingInputValues",
optional_params: dict,
headers: dict,
) -> dict:
"""
Transform embedding request for Cohere models on SageMaker
"""
if isinstance(input, str):
input_list: List[str] = [input]
elif isinstance(input, list):
if input and (isinstance(input[0], list) or isinstance(input[0], int)):
raise ValueError("Input must be a list of strings")
input_list = cast(List[str], input)
else:
input_list = [str(input)]
return dict(
BedrockCohereEmbeddingConfig()._transform_request(
model=model,
input=input_list,
inference_params=optional_params,
)
)
def transform_embedding_response(
self,
model: str,
raw_response: Response,
model_response: "EmbeddingResponse",
logging_obj: Any,
api_key: Optional[str] = None,
request_data: dict = {},
optional_params: dict = {},
litellm_params: dict = {},
) -> "EmbeddingResponse":
"""
Transform embedding response for Cohere models on SageMaker.
Uses `CohereEmbeddingConfig._populate_embedding_response` (not
`_transform_response`) so we do not log `post_call` a second time
— the SageMaker embedding handler already logs `post_call` before
invoking this transform.
"""
input_value = (
logging_obj.model_call_details.get("input")
or request_data.get("texts")
or request_data.get("images")
or []
)
if isinstance(input_value, str):
input_value = [input_value]
return CohereEmbeddingConfig()._populate_embedding_response(
response_json=raw_response.json(),
model_response=model_response,
model=model,
encoding=litellm.encoding,
input=input_value,
)
def validate_environment(
self,
headers: dict,
model: str,
messages: List[Any],
optional_params: dict,
litellm_params: dict,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> dict:
"""
Validate environment for SageMaker Cohere embeddings
"""
return {"Content-Type": "application/json"}

View file

@ -1,7 +1,7 @@
"""
Translate from OpenAI's `/v1/embeddings` to Sagemaker's `/invoke`
In the Huggingface TGI format.
In the Huggingface TGI format.
"""
from typing import TYPE_CHECKING, Any, List, Optional, Union
@ -11,12 +11,13 @@ if TYPE_CHECKING:
from httpx._models import Headers, Response
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.types.utils import Usage, EmbeddingResponse
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
from litellm.llms.voyage.embedding.transformation import VoyageEmbeddingConfig
from litellm.types.utils import EmbeddingResponse, Usage
from ..common_utils import SagemakerError
from .cohere_transformation import SagemakerCohereEmbeddingConfig
class SagemakerEmbeddingConfig(BaseEmbeddingConfig):
@ -38,17 +39,20 @@ class SagemakerEmbeddingConfig(BaseEmbeddingConfig):
Returns:
Appropriate embedding config instance
"""
if "voyage" in model.lower():
model_lower = model.lower()
if "voyage" in model_lower:
return VoyageEmbeddingConfig()
else:
return cls()
if "cohere" in model_lower:
return SagemakerCohereEmbeddingConfig()
return cls()
def get_supported_openai_params(self, model: str) -> List[str]:
# Check if this is an embedding model
if "voyage" in model.lower():
model_lower = model.lower()
if "voyage" in model_lower:
return VoyageEmbeddingConfig().get_supported_openai_params(model)
else:
return []
if "cohere" in model_lower:
return SagemakerCohereEmbeddingConfig().get_supported_openai_params(model)
return []
def map_openai_params(
self,

View file

@ -207,7 +207,7 @@ def resolve_resource_group(sources: List[Source]) -> Optional[str]:
def _parse_service_key_once(
service_key: Optional[Union[str, dict]]
service_key: Optional[Union[str, dict]],
) -> Optional[Dict[str, Any]]:
"""
Pre-parse service_key if it's a string to avoid repeated JSON parsing.

View file

@ -14,7 +14,6 @@ from ...openai_like.chat.transformation import OpenAIGPTConfig
from ..utils import SnowflakeBaseConfig
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj

View file

@ -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 TogetherAI is openai-compatible.

View file

@ -1,5 +1,5 @@
"""
Support for OpenAI's `/v1/embeddings` endpoint.
Support for OpenAI's `/v1/embeddings` endpoint.
Calls done in OpenAI/openai.py as TogetherAI is openai-compatible.

View file

@ -1,5 +1,5 @@
"""
Transformation logic from Cohere's /v1/rerank format to Together AI's `/v1/rerank` format.
Transformation logic from Cohere's /v1/rerank format to Together AI's `/v1/rerank` format.
Why separate file? Make it easy to see how transformation works
"""

View file

@ -1,5 +1,5 @@
"""
Transformation logic for context caching.
Transformation logic for context caching.
Why separate file? Make it easy to see how transformation works
"""
@ -19,7 +19,7 @@ from ..gemini.transformation import (
def get_first_continuous_block_idx(
filtered_messages: List[Tuple[int, AllMessageValues]] # (idx, message)
filtered_messages: List[Tuple[int, AllMessageValues]], # (idx, message)
) -> int:
"""
Find the array index that ends the first continuous sequence of message blocks.

View file

@ -1073,16 +1073,14 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915
contents.append(ContentType(role="user", parts=tool_call_responses))
if len(contents) == 0:
verbose_logger.warning(
"""
verbose_logger.warning("""
No contents in messages. Contents are required. See
https://cloud.google.com/vertex-ai/docs/reference/rest/v1/projects.locations.publishers.models/generateContent#request-body.
If the original request did not comply to OpenAI API requirements it should have failed by now,
but LiteLLM does not check for missing messages.
Setting an empty content to prevent an 400 error.
Relevant Issue - https://github.com/BerriAI/litellm/issues/9733
"""
)
""")
contents.append(ContentType(role="user", parts=[PartType(text=" ")]))
return contents
except Exception as e:

View file

@ -1,5 +1,5 @@
"""
Transformation logic from OpenAI /v1/embeddings format to Google AI Studio /batchEmbedContents format.
Transformation logic from OpenAI /v1/embeddings format to Google AI Studio /batchEmbedContents format.
Why separate file? Make it easy to see how transformation works
"""

View file

@ -139,7 +139,7 @@ class VertexTextToSpeechAPI(VertexLLM):
########## End of logging ############
####### Send the request ###################
if _is_async is True:
return self.async_audio_speech( # type:ignore
return self.async_audio_speech( # type: ignore
logging_obj=logging_obj, url=url, headers=headers, request=request
)
sync_handler = _get_httpx_client()

View file

@ -1,5 +1,5 @@
"""
Translates from OpenAI's `/v1/chat/completions` to the VLLM sdk `llm.generate`.
Translates from OpenAI's `/v1/chat/completions` to the VLLM sdk `llm.generate`.
NOT RECOMMENDED FOR PRODUCTION USE. Use `hosted_vllm/` instead.
"""

Some files were not shown because too many files have changed in this diff Show more