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

Co-authored-by: Mateo Wang <mateo-berri@users.noreply.github.com>
This commit is contained in:
Cursor Agent 2026-05-22 00:35:42 +00:00
commit ae5ee8f9f3
No known key found for this signature in database
334 changed files with 23636 additions and 3305 deletions

View file

@ -158,6 +158,8 @@ jobs:
CHOCOLATEY_CONFIRM_ALL: "true"
- run:
name: Install Dependencies
environment:
UV_HTTP_TIMEOUT: "300"
command: |
$installer = Join-Path $env:TEMP "uv-install.ps1"
Invoke-WebRequest -Uri https://astral.sh/uv/0.10.9/install.ps1 -OutFile $installer

View file

@ -215,8 +215,10 @@ jobs:
tests/proxy_unit_tests/test_models_fallback_endpoint.py
tests/proxy_unit_tests/test_google_endpoint_routing.py
tests/proxy_unit_tests/test_google_gemini_proxy_request.py
tests/proxy_unit_tests/test_gemini_agents_endpoints.py
tests/proxy_unit_tests/test_get_favicon.py
tests/proxy_unit_tests/test_get_image.py
tests/proxy_unit_tests/test_reducto_ocr_route.py
tests/proxy_unit_tests/test_ui_path_detection.py
tests/proxy_unit_tests/test_prompt_test_endpoint.py
tests/proxy_unit_tests/test_check_batch_cost.py

View file

@ -0,0 +1,34 @@
name: "Unit Tests: Proxy Management-Endpoint Behavior Pinning"
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_branch
- "litellm_**"
permissions:
contents: read
id-token: write
pull-requests: write
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs:
proxy-mgmt-behavior:
uses: ./.github/workflows/_test-unit-services-base.yml
with:
test-path: tests/proxy_behavior
# workers=0 (no xdist): the world seed is a single shared Postgres
# state — two xdist workers both call seed_world() and race on the
# ``behavior-pin-budget`` row, producing UniqueViolation + cascading
# missing-membership FK failures. The whole suite is ~7s sequentially,
# so the cost of disabling parallelism here is negligible.
workers: 0
reruns: 0
enable-postgres: true
artifact-name: proxy-mgmt-behavior
timeout-minutes: 15

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 \

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-proxy-extras"
version = "0.4.72"
version = "0.4.73"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.4.72"
version = "0.4.73"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-proxy-extras==",

View file

@ -225,6 +225,10 @@ use_chat_completions_url_for_anthropic_messages: bool = bool(
route_all_chat_openai_to_responses: bool = (
os.getenv("LITELLM_ROUTE_ALL_CHAT_OPENAI_TO_RESPONSES", "false").lower() == "true"
) # When True, routes all OpenAI /chat/completions requests through the Responses API bridge
use_legacy_interactions_schema: bool = (
os.getenv("LITELLM_USE_LEGACY_INTERACTIONS_SCHEMA", "false").lower() == "true"
) # When True, sends Api-Revision: 2026-05-07 to Google so responses use the legacy `outputs`
# schema instead of the new `steps` schema. Remove this flag after June 8, 2026.
retry = True
### AUTH ###
api_key: Optional[str] = None
@ -409,6 +413,12 @@ internal_user_budget_duration: Optional[str] = None
tag_budget_config: Optional[Dict[str, "BudgetConfig"]] = None
max_end_user_budget: Optional[float] = None
max_end_user_budget_id: Optional[str] = None
# When True, end-user IDs extracted from requests are validated against
# LiteLLM_EndUserTable / LiteLLM_UserTable. Values that do not resolve to a
# known row are dropped before reaching spend logs. Defaults to False for
# backwards compatibility — arbitrary client-supplied identifiers still
# pass through unchanged.
validate_end_user_id_in_db: bool = False
disable_end_user_cost_tracking: Optional[bool] = None
disable_end_user_cost_tracking_prometheus_only: Optional[bool] = None
enable_end_user_cost_tracking_prometheus_only: Optional[bool] = None
@ -632,6 +642,7 @@ minimax_models: Set = set()
aws_polly_models: Set = set()
gigachat_models: Set = set()
llamagate_models: Set = set()
reducto_models: Set = set()
bedrock_mantle_models: Set = set()
@ -899,6 +910,8 @@ def add_known_models(model_cost_map: Optional[Dict] = None):
gigachat_models.add(key)
elif value.get("litellm_provider") == "llamagate":
llamagate_models.add(key)
elif value.get("litellm_provider") == "reducto":
reducto_models.add(key)
elif value.get("litellm_provider") == "bedrock_mantle":
bedrock_mantle_models.add(key)
@ -1010,6 +1023,7 @@ model_list = list(
| ovhcloud_models
| lemonade_models
| docker_model_runner_models
| reducto_models
| bedrock_mantle_models
| set(clarifai_models)
)
@ -1116,6 +1130,7 @@ models_by_provider: dict = {
"aws_polly": aws_polly_models,
"gigachat": gigachat_models,
"llamagate": llamagate_models,
"reducto": reducto_models,
"bedrock_mantle": bedrock_mantle_models,
}
@ -1288,6 +1303,18 @@ from .responses.main import *
# Interactions API is available as litellm.interactions module
# Usage: litellm.interactions.create(), litellm.interactions.get(), etc.
from . import interactions
from .interactions.agents.main import (
acreate as acreate_agent,
create as create_agent,
alist as alist_agents,
list as list_agents,
aget as aget_agent,
get as get_agent,
adelete as adelete_agent,
delete as delete_agent,
alist_versions as alist_agent_versions,
list_versions as list_agent_versions,
)
from .skills.main import (
create_skill,
acreate_skill,

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

@ -30,6 +30,11 @@ from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
from litellm.llms.base_llm.bridges.completion_transformation import (
CompletionTransformationBridge,
)
from litellm.responses.sse_output_recovery import (
parse_sse_json_chunk,
record_output_item_chunk,
record_output_text_chunk,
)
from litellm.types.llms.openai import (
ChatCompletionAnnotation,
ChatCompletionReasoningItem,
@ -97,7 +102,7 @@ def _build_reasoning_item(
def _reasoning_item_to_response_input(
r_item: Union[ChatCompletionReasoningItem, Dict[str, Any]]
r_item: Union[ChatCompletionReasoningItem, Dict[str, Any]],
) -> Dict[str, Any]:
"""Convert a stored ChatCompletionReasoningItem back to a Responses API input item."""
r_input: Dict[str, Any] = {
@ -601,6 +606,79 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
return choices
@classmethod
def _extract_output_from_completed_event(
cls, parsed_chunk: Dict[str, Any]
) -> Optional[List[Dict[str, Any]]]:
response_payload = parsed_chunk.get("response")
if not isinstance(response_payload, dict):
return None
response_output = response_payload.get("output")
if not isinstance(response_output, list) or len(response_output) == 0:
return None
return cast(List[Dict[str, Any]], response_output)
@classmethod
def _recover_output_items_from_raw_sse(
cls, raw_sse: Optional[str]
) -> List[Dict[str, Any]]:
if not raw_sse or not isinstance(raw_sse, str):
return []
recovered_output_items: Dict[int, Dict[str, Any]] = {}
recovered_text_only_items: Dict[int, Dict[str, Any]] = {}
for chunk in raw_sse.splitlines():
parsed_chunk = parse_sse_json_chunk(chunk)
if parsed_chunk is None:
continue
event_type = parsed_chunk.get("type")
if event_type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED:
recovered_output = cls._extract_output_from_completed_event(
parsed_chunk
)
if recovered_output is not None:
return recovered_output
continue
if event_type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE:
record_output_item_chunk(
parsed_chunk=parsed_chunk,
output_items=recovered_output_items,
)
continue
if event_type == ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE:
record_output_text_chunk(
parsed_chunk=parsed_chunk,
output_items=recovered_output_items,
text_only_items=recovered_text_only_items,
)
continue
# Merge text-only items into the recovered output items. Real
# OUTPUT_ITEM_DONE events take precedence at any given output_index,
# but text-only items at indices without a matching OUTPUT_ITEM_DONE
# must still be preserved (e.g. multi-output responses where some
# indices only emitted OUTPUT_TEXT_DONE).
merged_items: Dict[int, Dict[str, Any]] = {**recovered_text_only_items}
merged_items.update(recovered_output_items)
if merged_items:
return [item for _, item in sorted(merged_items.items())]
return []
@classmethod
def _recover_output_items_from_logging(
cls, logging_obj: "LiteLLMLoggingObj"
) -> List[Dict[str, Any]]:
model_call_details = getattr(logging_obj, "model_call_details", {}) or {}
original_response = model_call_details.get("original_response")
return cls._recover_output_items_from_raw_sse(original_response)
def transform_response( # noqa: PLR0915
self,
model: str,
@ -625,9 +703,22 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
if raw_response.error is not None:
raise ValueError(f"Error in response: {raw_response.error}")
output_items = raw_response.output
if len(output_items) == 0:
recovered_output_items = self._recover_output_items_from_logging(
logging_obj
)
if recovered_output_items:
output_items = cast(Any, recovered_output_items)
raw_response.output = cast(Any, recovered_output_items)
verbose_logger.warning(
"Recovered empty Responses API output from raw SSE for model=%s",
model,
)
# Convert response output to choices using the static helper
choices = self._convert_response_output_to_choices(
output_items=raw_response.output,
output_items=output_items,
handle_raw_dict_callback=self._handle_raw_dict_response_item,
)
@ -641,7 +732,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
)
else:
raise ValueError(
f"Unknown items in responses API response: {raw_response.output}"
f"Unknown items in responses API response: {output_items}"
)
setattr(model_response, "choices", choices)
@ -1237,7 +1328,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
raise ValueError(
f"Chat provider: Invalid function argument delta {parsed_chunk}"
)
elif event_type == "response.output_item.done":
elif event_type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE:
# New output item added
output_item = parsed_chunk.get("item", {})
if output_item.get("type") == "function_call":

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

@ -1879,10 +1879,6 @@ def ocr_cost(
if response.usage_info is None:
raise ValueError("OCR response usage_info is None")
pages_processed = response.usage_info.pages_processed
if pages_processed is None:
raise ValueError("OCR response pages_processed is None")
try:
model_info: Optional[ModelInfo] = litellm.get_model_info(
model=model, custom_llm_provider=custom_llm_provider
@ -1890,9 +1886,49 @@ def ocr_cost(
except Exception:
model_info = None
ocr_cost_per_page: float = 0.0
credits = getattr(response.usage_info, "credits", None)
cost_per_credit = None
if model_info is not None:
ocr_cost_per_page = model_info.get("ocr_cost_per_page") or 0.0
cost_per_credit = model_info.get("ocr_cost_per_credit")
if credits is not None and cost_per_credit is not None:
return cost_per_credit * credits, 0.0
ocr_cost_per_page: Optional[float] = None
if model_info is not None:
ocr_cost_per_page = model_info.get("ocr_cost_per_page")
pages_processed = response.usage_info.pages_processed
if pages_processed is None:
if cost_per_credit is not None or ocr_cost_per_page is None:
# Surface missing usage data instead of silently under-reporting
# cost. The previous behavior raised ValueError; we now return 0.0
# for credit-priced or unpriced models, so log a warning to keep
# the regression visible to operators.
verbose_logger.warning(
"OCR cost: model=%s custom_llm_provider=%s response.usage_info."
"pages_processed is None and credits=%s; returning 0.0 cost.",
model,
custom_llm_provider,
credits,
)
return 0.0, 0.0
raise ValueError("OCR response pages_processed is None")
if ocr_cost_per_page is None:
# No per-page pricing configured. Either the model is on credit-based
# pricing (and credits weren't returned, so the credit branch above did
# not match) or the model has no OCR pricing entry at all. Surface a
# warning so that missing pricing entries are visible rather than
# silently producing zero cost for billable usage.
verbose_logger.warning(
"OCR cost: model=%s custom_llm_provider=%s reported "
"pages_processed=%s but no ocr_cost_per_page is configured; "
"returning 0.0 cost.",
model,
custom_llm_provider,
pages_processed,
)
return 0.0, 0.0
total_ocr_processing_cost: float = ocr_cost_per_page * pages_processed
return total_ocr_processing_cost, 0.0

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.
"""
@ -14,22 +14,38 @@ from litellm.integrations.custom_logger import CustomLogger
class CustomBatchLogger(CustomLogger):
preserve_events_added_during_flush = False
# Default cap on the in-memory log queue. Prevents unbounded memory growth
# if ``async_send_batch`` consistently fails (e.g. the destination is
# unreachable) and events are preserved across flush attempts. Subclasses
# may override by passing ``max_queue_size`` or by setting the attribute
# directly (see ``RubrikLogger`` for an example).
DEFAULT_MAX_QUEUE_SIZE = 50_000
def __init__(
self,
flush_lock: Optional[asyncio.Lock] = None,
batch_size: Optional[int] = None,
flush_interval: Optional[int] = None,
max_queue_size: Optional[int] = None,
**kwargs,
) -> None:
"""
Args:
flush_lock (Optional[asyncio.Lock], optional): Lock to use when flushing the queue. Defaults to None. Only used for custom loggers that do batching
max_queue_size (Optional[int], optional): Maximum number of events to retain in ``log_queue``. When the limit is exceeded (e.g. because the send destination is unreachable and events are preserved for retry), the oldest events are dropped. Defaults to ``DEFAULT_MAX_QUEUE_SIZE``.
"""
self.log_queue: List = []
self.flush_interval = flush_interval or litellm.DEFAULT_FLUSH_INTERVAL_SECONDS
self.batch_size: int = batch_size or litellm.DEFAULT_BATCH_SIZE
self.last_flush_time = time.time()
self.flush_lock = flush_lock
self.max_queue_size: int = (
max_queue_size
if max_queue_size is not None
else self.DEFAULT_MAX_QUEUE_SIZE
)
super().__init__(**kwargs)
@ -47,11 +63,40 @@ class CustomBatchLogger(CustomLogger):
async with self.flush_lock:
if self.log_queue:
log_queue_length = len(self.log_queue)
verbose_logger.debug(
"CustomLogger: Flushing batch of %s events", len(self.log_queue)
)
await self.async_send_batch()
self.log_queue.clear()
try:
await self.async_send_batch()
except Exception:
# If the underlying batch send raised, do NOT drop the
# in-flight events. They will be retried on the next flush.
# Most existing async_send_batch implementations swallow
# their own errors, so this only affects loggers that opt
# in to surfacing failures (e.g. Rubrik).
verbose_logger.exception(
"CustomLogger: async_send_batch raised; preserving "
"%s events in queue for retry",
log_queue_length,
)
# Guard against unbounded queue growth if the destination
# is persistently unreachable. Drop the oldest events
# beyond ``max_queue_size``.
overflow = len(self.log_queue) - self.max_queue_size
if overflow > 0:
del self.log_queue[:overflow]
verbose_logger.warning(
"CustomLogger: log queue exceeded max_queue_size=%s; "
"dropped %s oldest events.",
self.max_queue_size,
overflow,
)
return
if self.preserve_events_added_during_flush:
del self.log_queue[:log_queue_length]
else:
self.log_queue.clear()
self.last_flush_time = time.time()
async def async_send_batch(self, *args, **kwargs):

View file

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

View file

@ -673,6 +673,15 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
if parent_otel_span is not None:
parent_otel_span.set_status(Status(StatusCode.ERROR))
# Stamp team attributes onto the SERVER (root) span too, so the
# trace root is team-filterable on the failure path like the
# child exception span below.
self._set_team_attributes_on_span(
span=parent_otel_span,
team_id=user_api_key_dict.team_id,
team_alias=user_api_key_dict.team_alias,
)
# Stamp structured error attrs on the SERVER span itself; the
# failure path otherwise only sets its status (_handle_failure
# records on the litellm_request child span). Inline import:
@ -709,12 +718,65 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
key="exception",
value=str(original_exception),
)
self._set_team_attributes_on_span(
span=exception_logging_span,
team_id=user_api_key_dict.team_id,
team_alias=user_api_key_dict.team_alias,
)
exception_logging_span.set_status(Status(StatusCode.ERROR))
exception_logging_span.end(end_time=self._to_ns(datetime.now()))
# Emit guardrail spans for any guardrail invocations that
# ran during this request. _handle_failure typically does this,
# but for pre-call guardrail blocks the standard_logging_object
# may not carry guardrail_information by the time _handle_failure
# fires (the data lives only in request_data["metadata"]). Pull
# directly from request_data so the span is recorded either way;
# _emit_once dedupes if _handle_failure already emitted it.
self._emit_guardrail_spans_from_request_data(
request_data=request_data,
parent_span=parent_otel_span,
)
# End Parent OTEL Sspan
parent_otel_span.end(end_time=self._to_ns(datetime.now()))
def _emit_guardrail_spans_from_request_data(
self,
request_data: dict,
parent_span: Optional[Any],
) -> None:
"""Emit ``guardrail`` spans from ``request_data["metadata"]
["standard_logging_guardrail_information"]``.
Routed through ``_create_guardrail_span`` so the dedupe state in
``_otel_internal`` is honoured — if ``_handle_failure`` already
emitted these spans for the same kwargs, this is a no-op.
"""
from opentelemetry import trace as _trace
metadata = (request_data or {}).get("metadata") or {}
guardrail_information = metadata.get("standard_logging_guardrail_information")
if not guardrail_information:
return
# _create_guardrail_span reads guardrail_information from
# kwargs["standard_logging_object"] and shares its dedupe state via
# kwargs["litellm_params"]["metadata"]["_otel_internal"]. Pass the
# SAME metadata dict the proxy populated so _handle_failure and
# this hook see the same dedupe markers.
kwargs: Dict[str, Any] = {
"litellm_params": {"metadata": metadata},
"standard_logging_object": {
"guardrail_information": guardrail_information,
"metadata": metadata,
},
}
context = (
_trace.set_span_in_context(parent_span) if parent_span is not None else None
)
self._create_guardrail_span(kwargs=kwargs, context=context)
async def async_post_call_success_hook(
self,
data: dict,
@ -1012,6 +1074,10 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
):
parent_span.end(end_time=self._to_ns(end_time))
# Stamp team attributes onto the SERVER (root) span before it is
# closed, so the trace root carries them like every child span.
self._set_team_attributes_on_proxy_span_from_kwargs(kwargs)
# close the proxy span explicitly from kwargs metadata
# after all child spans (litellm_request, guardrail, raw_request)
# have been fully recorded and exported.
@ -1070,8 +1136,70 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
)
raw_span.set_status(Status(StatusCode.OK))
self.set_raw_request_attributes(raw_span, kwargs, response_obj)
self._set_team_attributes_from_kwargs(raw_span, kwargs)
raw_span.end(end_time=self._to_ns(end_time))
def _set_team_attributes_on_span(
self,
span: Span,
team_id: Optional[str],
team_alias: Optional[str],
) -> None:
"""Stamp team_id / team_alias onto a span so every child span of a
litellm_request trace carries them, not just the root span.
Empty strings are treated as absent: a request made with the master
key or a team-less virtual key carries ``user_api_key_team_id=""``
in ``standard_logging_object.metadata``; propagating that to every
span only adds noise that makes traces look mis-instrumented.
"""
if team_id:
self.safe_set_attribute(
span=span,
key="metadata.user_api_key_team_id",
value=team_id,
)
if team_alias:
self.safe_set_attribute(
span=span,
key="metadata.user_api_key_team_alias",
value=team_alias,
)
def _set_team_attributes_from_kwargs(self, span: Span, kwargs: dict) -> None:
"""Pull team_id / team_alias from the standard logging metadata in kwargs and stamp them onto span."""
std_log = kwargs.get("standard_logging_object")
md: dict = {}
if isinstance(std_log, dict):
md = std_log.get("metadata") or {}
elif std_log is not None:
md = getattr(std_log, "metadata", None) or {}
self._set_team_attributes_on_span(
span=span,
team_id=md.get("user_api_key_team_id"),
team_alias=md.get("user_api_key_team_alias"),
)
def _set_team_attributes_on_proxy_span_from_kwargs(self, kwargs: dict) -> None:
"""Stamp team attributes onto the proxy SERVER (root) span so the
trace root is filterable by team, not just its children. The root
span is created in auth before the team is resolved and is
otherwise only closed (never re-attributed) on the success path.
Guarded to the LiteLLM-created proxy span (by name + recording) so
externally provided parent spans are never mutated.
"""
litellm_params = kwargs.get("litellm_params") or {}
metadata = litellm_params.get("metadata") or {}
proxy_span = metadata.get("litellm_parent_otel_span")
if (
proxy_span is not None
and getattr(proxy_span, "name", None) == LITELLM_PROXY_REQUEST_SPAN_NAME
and hasattr(proxy_span, "is_recording")
and proxy_span.is_recording()
):
self._set_team_attributes_from_kwargs(proxy_span, kwargs)
def _record_metrics(self, kwargs, response_obj, start_time, end_time):
duration_s = (end_time - start_time).total_seconds()
params = kwargs.get("litellm_params") or {}
@ -1531,12 +1659,45 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
"masked_entity_count", safe_dumps(masked_entity_count)
)
guardrail_response = guardrail_information.get("guardrail_response")
if guardrail_response is not None:
guardrail_span.set_attribute(
"guardrail_response", safe_dumps(guardrail_response)
)
# Surface guardrail_status (success / guardrail_intervened /
# guardrail_failed_to_respond / not_run) as a top-level span
# attribute so trace backends can filter on it without parsing
# guardrail_response.
self.safe_set_attribute(
span=guardrail_span,
key="guardrail_response",
value=guardrail_information.get("guardrail_response"),
key="guardrail_status",
value=guardrail_information.get("guardrail_status"),
)
# Provider's raw top-level action (e.g. Bedrock's
# ``GUARDRAIL_INTERVENED`` / ``NONE``). Populated by the provider
# hook onto StandardLoggingGuardrailInformation so this integration
# stays provider-agnostic — we only read a normalised string.
guardrail_action = guardrail_information.get("guardrail_action")
if guardrail_action:
guardrail_span.set_attribute("guardrail_action", guardrail_action)
# The provider hook (e.g. Bedrock) extracts violation_categories
# from the raw response BEFORE redaction and stamps them onto
# StandardLoggingGuardrailInformation. Surfacing them here as a
# queryable attribute lets dashboards group by violation category
# without parsing the redacted guardrail_response blob.
violation_categories = guardrail_information.get("violation_categories")
if violation_categories:
# OTel sequence attributes must be homogeneous primitives;
# serialise to JSON once so set_attribute never coerces.
guardrail_span.set_attribute(
"guardrail_violation_categories", safe_dumps(violation_categories)
)
self._set_team_attributes_from_kwargs(guardrail_span, kwargs)
guardrail_span.end(end_time=self._to_ns(end_time_datetime))
def _handle_failure(self, kwargs, response_obj, start_time, end_time):

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

@ -0,0 +1,605 @@
"""Rubrik LiteLLM Plugin for tool blocking and batch logging."""
import asyncio
import os
import random
import time
import urllib.parse
import uuid
from collections import Counter
from typing import TYPE_CHECKING, Any, Literal, Optional
import httpx
from litellm._logging import verbose_logger
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
ModifyResponseException,
)
from litellm.litellm_core_utils.core_helpers import safe_deep_copy
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import (
ChatCompletionMessageToolCall,
Function,
GenericGuardrailAPIInputs,
StandardLoggingPayload,
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import (
Logging as LiteLLMLoggingObj,
)
_ENDPOINT_ANTHROPIC_MESSAGES = "/v1/messages"
_WEBHOOK_PATH_TOOL_BLOCKING = "/v1/after_completion/openai/v1"
_WEBHOOK_PATH_LOGGING_BATCH = "/v1/litellm/batch"
_MAX_QUEUE_SIZE = 10_000
_DROP_WARNING_INTERVAL_SECONDS = 60.0
class _MalformedToolBlockingResponseError(Exception):
"""Raised when the tool blocking service returns a structurally invalid
response (e.g. empty ``choices``).
Distinct from transient network/HTTP errors so callers can surface a
louder, misconfiguration-style log instead of treating it as a routine
fail-open.
"""
class RubrikLogger(CustomGuardrail, CustomBatchLogger):
def __init__(
self,
api_key: str | None = None,
api_base: str | None = None,
**kwargs,
):
self.flush_lock = asyncio.Lock()
kwargs.setdefault("guardrail_name", "rubrik")
# `initialize_guardrail` always passes these kwargs explicitly, with
# value `None` when the user omits `mode` / `default_on` from the
# guardrail config. Coerce None (omitted) to the desired default
# while preserving any explicit value the caller did set --
# in particular `default_on=False` if the user wants the guardrail
# off by default.
kwargs["event_hook"] = kwargs.get("event_hook") or GuardrailEventHooks.post_call
if kwargs.get("default_on") is None:
kwargs["default_on"] = True
super().__init__(
flush_lock=self.flush_lock,
**kwargs,
)
verbose_logger.debug("initializing rubrik logger")
self.sampling_rate = 1.0
rbrk_sampling_rate = os.getenv("RUBRIK_SAMPLING_RATE")
if rbrk_sampling_rate is not None:
try:
parsed_rate = float(rbrk_sampling_rate.strip())
self.sampling_rate = max(0.0, min(1.0, parsed_rate))
if parsed_rate != self.sampling_rate:
verbose_logger.warning(
f"RUBRIK_SAMPLING_RATE={parsed_rate} clamped to "
f"{self.sampling_rate}"
)
except ValueError:
verbose_logger.warning(
f"Invalid RUBRIK_SAMPLING_RATE: {rbrk_sampling_rate!r}, using 1.0"
)
self.key = api_key or os.getenv("RUBRIK_API_KEY")
if not self.key:
verbose_logger.warning(
"Rubrik: No API key configured. Requests will be unauthenticated."
)
_batch_size = os.getenv("RUBRIK_BATCH_SIZE")
if _batch_size:
try:
self.batch_size = int(_batch_size)
except ValueError:
verbose_logger.warning(
f"Invalid RUBRIK_BATCH_SIZE: {_batch_size!r}, using default"
)
# Cap the in-memory retry queue so a Rubrik webhook outage cannot let
# authenticated traffic accumulate prompt/response payloads until the
# proxy runs out of memory. Once the cap is reached, oldest events are
# dropped to make room for fresh ones (drop-oldest backpressure).
self.max_queue_size = _MAX_QUEUE_SIZE
self._dropped_since_warning = 0
self._last_drop_warning_time = 0.0
_webhook_url = api_base or os.getenv("RUBRIK_WEBHOOK_URL")
if _webhook_url is None:
raise ValueError(
"Rubrik webhook URL not configured. "
"Set RUBRIK_WEBHOOK_URL or pass api_base."
)
_webhook_url = _webhook_url.rstrip("/").removesuffix("/v1")
self.tool_blocking_endpoint = f"{_webhook_url}{_WEBHOOK_PATH_TOOL_BLOCKING}"
self.logging_endpoint = f"{_webhook_url}{_WEBHOOK_PATH_LOGGING_BATCH}"
self.async_httpx_client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.LoggingCallback
)
self.tool_blocking_client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.LoggingCallback,
params={"timeout": httpx.Timeout(5.0, connect=2.0)},
)
self._headers: dict[str, str] = {"Content-Type": "application/json"}
if self.key:
self._headers["Authorization"] = f"Bearer {self.key}"
# Periodic flush is started lazily on the first log event so that
# low-traffic deployments still get their batches drained even when the
# logger is instantiated outside a running event loop (sync init).
self._flush_task: Optional[asyncio.Task[Any]] = (
self._start_periodic_flush_task()
)
def _start_periodic_flush_task(self) -> Optional[asyncio.Task[Any]]:
"""Start the periodic flush task only when an event loop is already running."""
try:
loop = asyncio.get_running_loop()
except RuntimeError:
verbose_logger.debug(
"Rubrik logger init: no running event loop, "
"periodic flush will start on first log event."
)
return None
return loop.create_task(self.periodic_flush())
def _ensure_periodic_flush_task(self) -> None:
# Synchronous helper: in asyncio's cooperative model there is no await
# between the check and assignment, so two callers cannot race here.
if self._flush_task is None or self._flush_task.done():
self._flush_task = self._start_periodic_flush_task()
async def aclose(self):
"""Close the dedicated HTTP clients used by this logger."""
# Cancel the periodic flush task before closing the HTTP clients so
# the loop doesn't wake up and try to POST via a closed client.
if self._flush_task is not None and not self._flush_task.done():
self._flush_task.cancel()
try:
await self._flush_task
except (asyncio.CancelledError, Exception):
pass
self._flush_task = None
await self.tool_blocking_client.close()
await self.async_httpx_client.close()
# -- Guardrail hook --------------------------------------------------------
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional["LiteLLMLoggingObj"] = None,
) -> GenericGuardrailAPIInputs:
"""Validate tool calls against the blocking service (fail-open)."""
if input_type != "response":
return inputs
tool_calls = inputs.get("tool_calls")
if not tool_calls:
return inputs
try:
return await self._check_tool_calls(
inputs, tool_calls, request_data, logging_obj
)
except ModifyResponseException:
raise
except _MalformedToolBlockingResponseError as e:
# Distinct from transient errors: the service responded but the
# payload was structurally invalid, which usually indicates a
# misconfigured webhook or a breaking change in its response
# format. Log loudly so operators notice their tool-blocking
# policy is not actually being enforced.
verbose_logger.critical(
"Tool blocking service returned a malformed response: %s. "
"Tool calls are NOT being checked -- verify the webhook "
"configuration. Returning original response unchanged.",
e,
exc_info=True,
)
return inputs
except Exception as e:
verbose_logger.error(
f"Tool blocking hook failed: {e}. "
"Returning original response unchanged.",
exc_info=True,
)
return inputs
async def _check_tool_calls(
self,
inputs: GenericGuardrailAPIInputs,
tool_calls: Any,
request_data: dict,
logging_obj: Optional["LiteLLMLoggingObj"],
) -> GenericGuardrailAPIInputs:
"""Send tool calls to blocking service, raise if any are blocked."""
message_tool_calls = self._normalize_tool_calls(tool_calls)
call_details = (
getattr(logging_obj, "model_call_details", {}) if logging_obj else {}
)
response = request_data.get("response")
request_id = getattr(response, "id", None) if response else None
if logging_obj and not call_details:
verbose_logger.warning(
"Rubrik: logging_obj present but model_call_details is empty "
"-- request context will be missing"
)
response_data = self._build_tool_call_payload(message_tool_calls, request_id)
req_data = self._extract_request_data(call_details)
service_response = await self._post_to_tool_blocking_service(
response_data, req_data
)
blocked_explanation = self._extract_blocked_tools(
service_response, message_tool_calls
)
if blocked_explanation is not None:
model = self._resolve_model(request_data, call_details)
raise ModifyResponseException(
message=blocked_explanation,
model=model,
request_data=request_data,
guardrail_name=self.guardrail_name,
)
return inputs
@staticmethod
def _normalize_tool_calls(tool_calls: Any) -> list[ChatCompletionMessageToolCall]:
"""Convert tool_calls from inputs to ChatCompletionMessageToolCall objects."""
result = []
for tc in tool_calls:
if isinstance(tc, ChatCompletionMessageToolCall):
result.append(tc)
elif isinstance(tc, dict):
func = tc.get("function", {})
result.append(
ChatCompletionMessageToolCall(
id=tc.get("id", ""),
type=tc.get("type", "function"),
function=Function(
name=func.get("name", ""),
arguments=func.get("arguments", ""),
),
)
)
elif hasattr(tc, "id") and hasattr(tc, "function"):
result.append(
ChatCompletionMessageToolCall(
id=tc.id or "",
type=getattr(tc, "type", None) or "function",
function=tc.function,
)
)
else:
raise TypeError(
f"Cannot normalize tool_call of type {type(tc).__name__}"
)
return result
@staticmethod
def _build_tool_call_payload(
tool_calls: list[ChatCompletionMessageToolCall],
request_id: str | None,
) -> dict[str, Any]:
"""Build a full OpenAI ChatCompletion-format dict for the blocking service."""
return {
"id": request_id or f"chatcmpl-{uuid.uuid4()}",
"object": "chat.completion",
"created": int(time.time()),
"model": "",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": None,
"tool_calls": [
tc.model_dump(exclude_none=True) for tc in tool_calls
],
},
"finish_reason": "tool_calls",
}
],
}
@staticmethod
def _extract_request_data(call_details: dict[str, Any]) -> dict[str, Any]:
"""Extract original request data from model_call_details."""
if not call_details:
return {}
litellm_params = call_details.get("litellm_params", {}) or {}
return {
"messages": call_details.get("messages"),
"model": call_details.get("model"),
"proxy_server_request": RubrikLogger._sanitize_proxy_server_request(
litellm_params.get("proxy_server_request")
),
}
@staticmethod
def _sanitize_proxy_server_request(proxy_server_request: Any) -> Any:
"""Allowlist only routing fields (``url``, ``method``) when forwarding
``proxy_server_request`` to the external Rubrik webhook, dropping
inbound ``headers`` (Authorization, Cookie, x-api-key, ...) and the raw
request ``body`` so proxy credentials are not exfiltrated."""
if not isinstance(proxy_server_request, dict):
return proxy_server_request
return {
key: proxy_server_request[key]
for key in ("url", "method")
if key in proxy_server_request
}
@staticmethod
def _resolve_model(
request_data: dict[str, Any], call_details: dict[str, Any]
) -> str:
"""Get the model name for the ModifyResponseException."""
response = request_data.get("response")
if response and hasattr(response, "model"):
return response.model or "unknown"
return call_details.get("model", "unknown")
# -- Logging hooks ---------------------------------------------------------
async def _prepare_log_payload(
self, kwargs: dict, event_type: str
) -> StandardLoggingPayload | None:
"""Shared logic for success and failure logging."""
if random.random() > self.sampling_rate:
verbose_logger.debug(
f"Skipping Rubrik {event_type} logging "
f"(sampling_rate={self.sampling_rate})"
)
return None
# Deep-copy so mutations don't affect other callbacks sharing this object
standard_logging_payload: StandardLoggingPayload = safe_deep_copy(
kwargs["standard_logging_object"]
)
# For Anthropic /v1/messages requests, LiteLLM creates a separate
# ModelResponse (with a generated chatcmpl-* id) for logging, which
# differs from the original Anthropic msg-* id on the response dict.
# Normalize to litellm_call_id so that the logging and tool-blocking
# endpoints see the same request identifier.
litellm_params = kwargs.get("litellm_params", {}) or {}
proxy_request = litellm_params.get("proxy_server_request", {}) or {}
url_path = urllib.parse.urlparse(proxy_request.get("url", "")).path
if url_path.endswith(_ENDPOINT_ANTHROPIC_MESSAGES):
_litellm_call_id = kwargs.get("litellm_call_id")
if _litellm_call_id:
standard_logging_payload["id"] = _litellm_call_id # type: ignore[literal-required]
if "system" in kwargs:
system_prompt_msg_list = kwargs["system"]
try:
if system_prompt_msg_list:
system_scaffold = {
"role": "system",
"content": system_prompt_msg_list,
}
if isinstance(standard_logging_payload["messages"], list):
standard_logging_payload["messages"].insert(0, system_scaffold)
elif isinstance(standard_logging_payload["messages"], (dict, str)):
standard_logging_payload["messages"] = [
system_scaffold,
standard_logging_payload["messages"],
]
except Exception as e:
verbose_logger.warning(
f"Rubrik: failed to prepend system prompt: {e}",
exc_info=True,
)
return standard_logging_payload
async def _enqueue_log_event(self, kwargs: dict, event_type: str):
try:
self._ensure_periodic_flush_task()
payload = await self._prepare_log_payload(kwargs, event_type)
if payload is None:
return
self.log_queue.append(payload)
self._enforce_max_queue_size()
if len(self.log_queue) >= self.batch_size:
await self.flush_queue()
except Exception as e:
verbose_logger.error(
f"Rubrik {event_type} logging hook failed: {e}. "
"Skipping logging for this event.",
exc_info=True,
)
def _enforce_max_queue_size(self) -> None:
overflow = len(self.log_queue) - self.max_queue_size
if overflow <= 0:
return
del self.log_queue[:overflow]
self._dropped_since_warning += overflow
now = time.time()
if now - self._last_drop_warning_time >= _DROP_WARNING_INTERVAL_SECONDS:
verbose_logger.warning(
"Rubrik: log queue exceeded max_queue_size=%s; dropped %s "
"oldest events since the last warning. The Rubrik webhook may "
"be unhealthy or undersized for current traffic.",
self.max_queue_size,
self._dropped_since_warning,
)
self._dropped_since_warning = 0
self._last_drop_warning_time = now
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
await self._enqueue_log_event(kwargs, "success")
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
await self._enqueue_log_event(kwargs, "failure")
# -- Batch logging ---------------------------------------------------------
async def _log_batch_to_rubrik(self, data):
# NOTE: this method intentionally re-raises on failure so the parent
# CustomBatchLogger.flush_queue keeps the unsent events in the queue
# for the next flush attempt instead of silently dropping them.
try:
response = await self.async_httpx_client.post(
url=self.logging_endpoint,
json=data,
headers=self._headers,
)
response.raise_for_status()
except httpx.HTTPStatusError as e:
verbose_logger.exception(
f"Rubrik HTTP Error: {e.response.status_code} - {e.response.text}"
)
raise
except Exception:
verbose_logger.exception("Rubrik Layer Error")
raise
async def async_send_batch(self):
"""Handles sending batches of responses to Rubrik.
Note: the canonical flush path is :meth:`flush_queue`, which takes a
single snapshot used for both sending and queue draining. This method
is kept for direct callers / tests; it intentionally does NOT remove
events from the queue.
"""
if not self.log_queue:
return
log_queue_snapshot = list(self.log_queue)
verbose_logger.debug(
"Rubrik: Flushing batch of %s events", len(log_queue_snapshot)
)
await self._log_batch_to_rubrik(
data=log_queue_snapshot,
)
async def flush_queue(self):
"""Snapshot, send, and drain in one consistent step.
Overrides the base implementation so the same snapshot drives both
the HTTP send and the queue truncation. This avoids the subtle
coupling where the base class captures `len(self.log_queue)`
separately from the snapshot taken inside `async_send_batch`,
which could otherwise drift in a future refactor and cause
duplicate deliveries to Rubrik.
"""
if self.flush_lock is None:
return
async with self.flush_lock:
if not self.log_queue:
return
snapshot = list(self.log_queue)
verbose_logger.debug("Rubrik: Flushing batch of %s events", len(snapshot))
try:
await self._log_batch_to_rubrik(data=snapshot)
except Exception:
# Already logged with traceback inside _log_batch_to_rubrik.
# Preserve the in-flight events for retry on the next flush.
return
del self.log_queue[: len(snapshot)]
self.last_flush_time = time.time()
# -- Tool blocking service -------------------------------------------------
async def _post_to_tool_blocking_service(
self,
response_data: dict[str, Any],
request_data: dict[str, Any],
) -> dict[str, Any]:
"""Post a payload to the tool blocking service and return the response.
Args:
response_data: The OpenAI-formatted response payload to send.
request_data: Original LLM request data to include alongside
the response for additional context. Empty dict if unavailable.
Raises:
Exception: If the service is unavailable or returns an error.
"""
envelope = {
"request": request_data,
"response": response_data,
}
verbose_logger.debug(
f"Sending request to tool blocking service: "
f"{self.tool_blocking_endpoint}"
)
http_response = await self.tool_blocking_client.post(
self.tool_blocking_endpoint,
json=envelope,
headers=self._headers,
)
http_response.raise_for_status()
result: dict[str, Any] = http_response.json()
return result
@staticmethod
def _extract_blocked_tools(
service_response: dict[str, Any],
all_tool_calls: list[ChatCompletionMessageToolCall],
) -> Optional[str]:
"""Return the blocking explanation if any tool calls were blocked.
Compares the service response (which contains only allowed tools) against
the full set of tool calls. Returns ``None`` if all tools are allowed, or
the explanation string (prefixed with newlines) otherwise.
Expects service_response in OpenAI chat completion format:
{"choices": [{"message": {"tool_calls": [...], "content": "..."}}]}
"""
choices = service_response.get("choices", [])
if not choices:
raise _MalformedToolBlockingResponseError(
"Tool blocking service returned empty response"
)
message = choices[0].get("message", {})
returned_tool_calls = message.get("tool_calls") or []
blocking_explanation = message.get("content", "")
allowed_id_counts: Counter = Counter(
tc["id"]
for tc in returned_tool_calls
if isinstance(tc, dict) and tc.get("id")
)
required_id_counts: Counter = Counter(tc.id for tc in all_tool_calls if tc.id)
all_allowed = len(returned_tool_calls) >= len(all_tool_calls) and all(
allowed_id_counts.get(tc_id, 0) >= count
for tc_id, count in required_id_counts.items()
)
if all_allowed:
return None
explanation = blocking_explanation or "Tool call blocked by policy."
return f"\n\n{explanation}"

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

@ -5,31 +5,40 @@ This module provides SDK methods for Google's Interactions API.
Usage:
import litellm
# Create an interaction with a model
response = litellm.interactions.create(
model="gemini-2.5-flash",
input="Hello, how are you?"
)
# Create an interaction with an agent
response = litellm.interactions.create(
agent="deep-research-pro-preview-12-2025",
input="Research the current state of cancer research"
)
# Async version
response = await litellm.interactions.acreate(...)
# Get an interaction
response = litellm.interactions.get(interaction_id="...")
# Delete an interaction
result = litellm.interactions.delete(interaction_id="...")
# Cancel an interaction
result = litellm.interactions.cancel(interaction_id="...")
# Create a managed agent on the provider side
result = litellm.interactions.agents.create(
name="waverunner",
custom_llm_provider="gemini",
api_key="...",
base_agent="gemini-2.5-flash",
instructions="You are a helpful assistant.",
)
Methods:
- create(): Sync create interaction
- acreate(): Async create interaction
@ -39,8 +48,12 @@ Methods:
- adelete(): Async delete interaction
- cancel(): Sync cancel interaction
- acancel(): Async cancel interaction
Sub-modules:
- agents: Provider-side agent creation (litellm.interactions.agents.create)
"""
from litellm.interactions import agents
from litellm.interactions.main import (
acancel,
acreate,
@ -65,4 +78,6 @@ __all__ = [
# Cancel
"cancel",
"acancel",
# Sub-modules
"agents",
]

View file

@ -0,0 +1,39 @@
"""
litellm.interactions.agents
Full CRUD SDK for provider-side managed agents (e.g. Gemini v1beta/agents).
litellm.interactions.agents.create(name=..., ...)
litellm.interactions.agents.list(api_key=...)
litellm.interactions.agents.get(name=..., ...)
litellm.interactions.agents.delete(name=..., ...)
litellm.interactions.agents.list_versions(name=..., ...)
Async counterparts: acreate, alist, aget, adelete, alist_versions
"""
from litellm.interactions.agents.main import (
acreate,
adelete,
aget,
alist,
alist_versions,
create,
delete,
get,
list,
list_versions,
)
__all__ = [
"create",
"acreate",
"list",
"alist",
"get",
"aget",
"delete",
"adelete",
"list_versions",
"alist_versions",
]

View file

@ -0,0 +1,478 @@
"""
HTTP handler for the Agents API.
Extends InteractionsHTTPHandler so that the shared HTTP infrastructure
(_handle_error, _sync_client, _async_client) is reused rather than
duplicated. BaseAgentsAPIConfig stays as pure transform code.
"""
from typing import Any, Coroutine, Dict, Optional, Union
import httpx
from litellm.constants import request_timeout
from litellm.interactions.http_handler import InteractionsHTTPHandler
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.agents.transformation import BaseAgentsAPIConfig
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.types.agents import (
AgentCreateResponse,
AgentDeleteResult,
AgentListResponse,
AgentVersionsResponse,
)
from litellm.types.router import GenericLiteLLMParams
class AgentsHTTPHandler(InteractionsHTTPHandler):
"""HTTP handler for Agents API CRUD requests."""
# ------------------------------------------------------------------ #
# CREATE #
# ------------------------------------------------------------------ #
def create_agent(
self,
agents_api_config: BaseAgentsAPIConfig,
name: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[HTTPHandler] = None,
_is_async: bool = False,
) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]:
if _is_async:
return self.async_create_agent(
agents_api_config=agents_api_config,
name=name,
litellm_params=litellm_params,
logging_obj=logging_obj,
extra_headers=extra_headers,
extra_body=extra_body,
timeout=timeout,
)
sync_httpx_client = self._sync_client(litellm_params, client)
headers = agents_api_config.validate_environment(
headers=extra_headers or {}, litellm_params=dict(litellm_params)
)
url = agents_api_config.get_complete_url(
api_base=litellm_params.get("api_base"),
litellm_params=dict(litellm_params),
)
data = agents_api_config.transform_create_request(
name=name, litellm_params=dict(litellm_params)
)
if extra_body:
data.update(extra_body)
logging_obj.pre_call(
input=name,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": url,
"headers": headers,
},
)
try:
response = sync_httpx_client.post(
url=url, headers=headers, json=data, timeout=timeout or request_timeout
)
except Exception as e:
raise self._handle_error(e=e, provider_config=agents_api_config)
logging_obj.post_call(
original_response=response.text,
additional_args={"complete_input_dict": data},
)
return agents_api_config.transform_create_response(
raw_response=response, name=name
)
async def async_create_agent(
self,
agents_api_config: BaseAgentsAPIConfig,
name: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[AsyncHTTPHandler] = None,
) -> AgentCreateResponse:
async_httpx_client = self._async_client(litellm_params, client)
headers = agents_api_config.validate_environment(
headers=extra_headers or {}, litellm_params=dict(litellm_params)
)
url = agents_api_config.get_complete_url(
api_base=litellm_params.get("api_base"),
litellm_params=dict(litellm_params),
)
data = agents_api_config.transform_create_request(
name=name, litellm_params=dict(litellm_params)
)
if extra_body:
data.update(extra_body)
logging_obj.pre_call(
input=name,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": url,
"headers": headers,
},
)
try:
response = await async_httpx_client.post(
url=url, headers=headers, json=data, timeout=timeout or request_timeout
)
except Exception as e:
raise self._handle_error(e=e, provider_config=agents_api_config)
logging_obj.post_call(
original_response=response.text,
additional_args={"complete_input_dict": data},
)
return agents_api_config.transform_create_response(
raw_response=response, name=name
)
# ------------------------------------------------------------------ #
# LIST #
# ------------------------------------------------------------------ #
def list_agents(
self,
agents_api_config: BaseAgentsAPIConfig,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[HTTPHandler] = None,
_is_async: bool = False,
) -> Union[AgentListResponse, Coroutine[Any, Any, AgentListResponse]]:
if _is_async:
return self.async_list_agents(
agents_api_config=agents_api_config,
litellm_params=litellm_params,
logging_obj=logging_obj,
extra_headers=extra_headers,
timeout=timeout,
)
sync_httpx_client = self._sync_client(litellm_params, client)
headers = agents_api_config.validate_environment(
headers=extra_headers or {}, litellm_params=dict(litellm_params)
)
url, params = agents_api_config.transform_list_request(
api_base=litellm_params.get("api_base"),
litellm_params=dict(litellm_params),
)
logging_obj.pre_call(
input="list_agents",
api_key="",
additional_args={"api_base": url, "headers": headers},
)
try:
response = sync_httpx_client.get(url=url, headers=headers, params=params)
except Exception as e:
raise self._handle_error(e=e, provider_config=agents_api_config)
logging_obj.post_call(original_response=response.text, additional_args={})
return agents_api_config.transform_list_response(raw_response=response)
async def async_list_agents(
self,
agents_api_config: BaseAgentsAPIConfig,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[AsyncHTTPHandler] = None,
) -> AgentListResponse:
async_httpx_client = self._async_client(litellm_params, client)
headers = agents_api_config.validate_environment(
headers=extra_headers or {}, litellm_params=dict(litellm_params)
)
url, params = agents_api_config.transform_list_request(
api_base=litellm_params.get("api_base"),
litellm_params=dict(litellm_params),
)
logging_obj.pre_call(
input="list_agents",
api_key="",
additional_args={"api_base": url, "headers": headers},
)
try:
response = await async_httpx_client.get(
url=url, headers=headers, params=params
)
except Exception as e:
raise self._handle_error(e=e, provider_config=agents_api_config)
logging_obj.post_call(original_response=response.text, additional_args={})
return agents_api_config.transform_list_response(raw_response=response)
# ------------------------------------------------------------------ #
# GET #
# ------------------------------------------------------------------ #
def get_agent(
self,
agents_api_config: BaseAgentsAPIConfig,
name: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[HTTPHandler] = None,
_is_async: bool = False,
) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]:
if _is_async:
return self.async_get_agent(
agents_api_config=agents_api_config,
name=name,
litellm_params=litellm_params,
logging_obj=logging_obj,
extra_headers=extra_headers,
timeout=timeout,
)
sync_httpx_client = self._sync_client(litellm_params, client)
headers = agents_api_config.validate_environment(
headers=extra_headers or {}, litellm_params=dict(litellm_params)
)
url, params = agents_api_config.transform_get_request(
name=name,
api_base=litellm_params.get("api_base"),
litellm_params=dict(litellm_params),
)
logging_obj.pre_call(
input=name,
api_key="",
additional_args={"api_base": url, "headers": headers},
)
try:
response = sync_httpx_client.get(url=url, headers=headers, params=params)
except Exception as e:
raise self._handle_error(e=e, provider_config=agents_api_config)
logging_obj.post_call(original_response=response.text, additional_args={})
return agents_api_config.transform_get_response(
raw_response=response, name=name
)
async def async_get_agent(
self,
agents_api_config: BaseAgentsAPIConfig,
name: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[AsyncHTTPHandler] = None,
) -> AgentCreateResponse:
async_httpx_client = self._async_client(litellm_params, client)
headers = agents_api_config.validate_environment(
headers=extra_headers or {}, litellm_params=dict(litellm_params)
)
url, params = agents_api_config.transform_get_request(
name=name,
api_base=litellm_params.get("api_base"),
litellm_params=dict(litellm_params),
)
logging_obj.pre_call(
input=name,
api_key="",
additional_args={"api_base": url, "headers": headers},
)
try:
response = await async_httpx_client.get(
url=url, headers=headers, params=params
)
except Exception as e:
raise self._handle_error(e=e, provider_config=agents_api_config)
logging_obj.post_call(original_response=response.text, additional_args={})
return agents_api_config.transform_get_response(
raw_response=response, name=name
)
# ------------------------------------------------------------------ #
# DELETE #
# ------------------------------------------------------------------ #
def delete_agent(
self,
agents_api_config: BaseAgentsAPIConfig,
name: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[HTTPHandler] = None,
_is_async: bool = False,
) -> Union[AgentDeleteResult, Coroutine[Any, Any, AgentDeleteResult]]:
if _is_async:
return self.async_delete_agent(
agents_api_config=agents_api_config,
name=name,
litellm_params=litellm_params,
logging_obj=logging_obj,
extra_headers=extra_headers,
timeout=timeout,
)
sync_httpx_client = self._sync_client(litellm_params, client)
headers = agents_api_config.validate_environment(
headers=extra_headers or {}, litellm_params=dict(litellm_params)
)
url = agents_api_config.transform_delete_request(
name=name,
api_base=litellm_params.get("api_base"),
litellm_params=dict(litellm_params),
)
logging_obj.pre_call(
input=name,
api_key="",
additional_args={"api_base": url, "headers": headers},
)
try:
response = sync_httpx_client.delete(
url=url, headers=headers, timeout=timeout or request_timeout
)
except Exception as e:
raise self._handle_error(e=e, provider_config=agents_api_config)
logging_obj.post_call(original_response=response.text, additional_args={})
return agents_api_config.transform_delete_response(
raw_response=response, name=name
)
async def async_delete_agent(
self,
agents_api_config: BaseAgentsAPIConfig,
name: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[AsyncHTTPHandler] = None,
) -> AgentDeleteResult:
async_httpx_client = self._async_client(litellm_params, client)
headers = agents_api_config.validate_environment(
headers=extra_headers or {}, litellm_params=dict(litellm_params)
)
url = agents_api_config.transform_delete_request(
name=name,
api_base=litellm_params.get("api_base"),
litellm_params=dict(litellm_params),
)
logging_obj.pre_call(
input=name,
api_key="",
additional_args={"api_base": url, "headers": headers},
)
try:
response = await async_httpx_client.delete(
url=url, headers=headers, timeout=timeout or request_timeout
)
except Exception as e:
raise self._handle_error(e=e, provider_config=agents_api_config)
logging_obj.post_call(original_response=response.text, additional_args={})
return agents_api_config.transform_delete_response(
raw_response=response, name=name
)
# ------------------------------------------------------------------ #
# LIST VERSIONS #
# ------------------------------------------------------------------ #
def list_agent_versions(
self,
agents_api_config: BaseAgentsAPIConfig,
name: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[HTTPHandler] = None,
_is_async: bool = False,
) -> Union[AgentVersionsResponse, Coroutine[Any, Any, AgentVersionsResponse]]:
if _is_async:
return self.async_list_agent_versions(
agents_api_config=agents_api_config,
name=name,
litellm_params=litellm_params,
logging_obj=logging_obj,
extra_headers=extra_headers,
timeout=timeout,
)
sync_httpx_client = self._sync_client(litellm_params, client)
headers = agents_api_config.validate_environment(
headers=extra_headers or {}, litellm_params=dict(litellm_params)
)
url, params = agents_api_config.transform_list_versions_request(
name=name,
api_base=litellm_params.get("api_base"),
litellm_params=dict(litellm_params),
)
logging_obj.pre_call(
input=name,
api_key="",
additional_args={"api_base": url, "headers": headers},
)
try:
response = sync_httpx_client.get(url=url, headers=headers, params=params)
except Exception as e:
raise self._handle_error(e=e, provider_config=agents_api_config)
logging_obj.post_call(original_response=response.text, additional_args={})
return agents_api_config.transform_list_versions_response(
raw_response=response, name=name
)
async def async_list_agent_versions(
self,
agents_api_config: BaseAgentsAPIConfig,
name: str,
litellm_params: GenericLiteLLMParams,
logging_obj: LiteLLMLoggingObj,
extra_headers: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[AsyncHTTPHandler] = None,
) -> AgentVersionsResponse:
async_httpx_client = self._async_client(litellm_params, client)
headers = agents_api_config.validate_environment(
headers=extra_headers or {}, litellm_params=dict(litellm_params)
)
url, params = agents_api_config.transform_list_versions_request(
name=name,
api_base=litellm_params.get("api_base"),
litellm_params=dict(litellm_params),
)
logging_obj.pre_call(
input=name,
api_key="",
additional_args={"api_base": url, "headers": headers},
)
try:
response = await async_httpx_client.get(
url=url, headers=headers, params=params
)
except Exception as e:
raise self._handle_error(e=e, provider_config=agents_api_config)
logging_obj.post_call(original_response=response.text, additional_args={})
return agents_api_config.transform_list_versions_response(
raw_response=response, name=name
)
agents_http_handler = AgentsHTTPHandler()

View file

@ -0,0 +1,522 @@
"""
LiteLLM Agents API - Main Module
Usage:
import litellm
# Create
response = litellm.interactions.agents.create(
name="waverunner",
custom_llm_provider="gemini",
api_key="...",
base_agent="gemini-2.5-flash",
instructions="You are a helpful assistant.",
)
# List
response = litellm.interactions.agents.list(api_key="...", custom_llm_provider="gemini")
# Get
response = litellm.interactions.agents.get(name="waverunner", api_key="...")
# Delete
result = litellm.interactions.agents.delete(name="waverunner", api_key="...")
# List versions
result = litellm.interactions.agents.list_versions(name="waverunner", api_key="...")
# Async versions: acreate, alist, aget, adelete, alist_versions
"""
import asyncio
import contextvars
from functools import partial
from typing import Any, Coroutine, Dict, Optional, Union
import httpx
import litellm
from litellm.interactions.agents.http_handler import agents_http_handler
from litellm.interactions.agents.utils import get_provider_agents_api_config
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.agents import (
AgentCreateResponse,
AgentDeleteResult,
AgentListResponse,
AgentVersionsResponse,
)
from litellm.types.interactions import InteractionEnvironment
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import client
# ------------------------------------------------------------------ #
# Shared helpers #
# ------------------------------------------------------------------ #
def _get_agents_api_config(custom_llm_provider: str):
config = get_provider_agents_api_config(custom_llm_provider)
if config is None:
raise litellm.BadRequestError(
message=(
f"Provider '{custom_llm_provider}' does not have a native "
"agents API. Use the proxy POST /v1/agents endpoint to store "
"agents locally."
),
model="",
llm_provider=custom_llm_provider,
)
return config
def _make_logging_obj(
kwargs: Dict[str, Any],
model: str,
custom_llm_provider: str,
call_type: str,
optional_params: Dict[str, Any],
) -> LiteLLMLoggingObj:
litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj") # type: ignore
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
optional_params=optional_params,
litellm_params={"litellm_call_id": litellm_call_id},
custom_llm_provider=custom_llm_provider,
)
return litellm_logging_obj
# ================================================================== #
# CREATE #
# ================================================================== #
@client
async def acreate(
name: str,
base_agent: Optional[str] = None,
instructions: Optional[str] = None,
base_environment: Optional[InteractionEnvironment] = None,
custom_llm_provider: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
**kwargs,
) -> AgentCreateResponse:
"""Async: Create a managed agent on the provider side."""
local_vars = locals()
try:
loop = asyncio.get_event_loop()
kwargs["acreate_agent"] = True
func = partial(
create,
name=name,
base_agent=base_agent,
instructions=instructions,
base_environment=base_environment,
custom_llm_provider=custom_llm_provider or "gemini",
extra_headers=extra_headers,
extra_body=extra_body,
timeout=timeout,
**kwargs,
)
ctx = contextvars.copy_context()
init_response = await loop.run_in_executor(None, partial(ctx.run, func))
if asyncio.iscoroutine(init_response):
return await init_response
return init_response
except Exception as e:
raise litellm.exception_type(
model=name,
custom_llm_provider=custom_llm_provider or "gemini",
original_exception=e,
completion_kwargs=local_vars,
extra_kwargs=kwargs,
)
@client
def create(
name: str,
base_agent: Optional[str] = None,
instructions: Optional[str] = None,
base_environment: Optional[InteractionEnvironment] = None,
custom_llm_provider: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
**kwargs,
) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]:
"""
Sync: Create a managed agent on the provider side.
Args:
name: Name for the agent (required).
base_agent: Base agent to derive from (e.g. "waverunner").
instructions: System instructions for the agent.
base_environment: Environment to fork from — an env_id string or a
dict like ``{"type": "remote", "sources": [...]}``.
custom_llm_provider: Provider to use, e.g. "gemini".
extra_headers: Additional HTTP headers.
extra_body: Additional request body fields.
timeout: Request timeout.
**kwargs: Forwarded to GenericLiteLLMParams (api_key, api_base, etc.).
"""
local_vars = locals()
custom_llm_provider = (
custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini"
)
try:
_is_async = kwargs.pop("acreate_agent", False) is True
if base_agent is not None:
kwargs["base_agent"] = base_agent
if instructions is not None:
kwargs["instructions"] = instructions
if base_environment is not None:
kwargs["base_environment"] = base_environment
kwargs.setdefault("custom_llm_provider", custom_llm_provider)
litellm_params = GenericLiteLLMParams(**kwargs)
logging_obj = _make_logging_obj(
kwargs, name, custom_llm_provider, "create_agent", {}
)
config = _get_agents_api_config(custom_llm_provider)
return agents_http_handler.create_agent(
agents_api_config=config,
name=name,
litellm_params=litellm_params,
logging_obj=logging_obj,
extra_headers=extra_headers,
extra_body=extra_body,
timeout=timeout,
_is_async=_is_async,
)
except Exception as e:
raise litellm.exception_type(
model=name,
custom_llm_provider=custom_llm_provider,
original_exception=e,
completion_kwargs=local_vars,
extra_kwargs=kwargs,
)
# ================================================================== #
# LIST #
# ================================================================== #
@client
async def alist(
custom_llm_provider: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
**kwargs,
) -> AgentListResponse:
"""Async: List all agents on the provider side."""
local_vars = locals()
try:
loop = asyncio.get_event_loop()
kwargs["alist_agents"] = True
func = partial(
list,
custom_llm_provider=custom_llm_provider or "gemini",
extra_headers=extra_headers,
timeout=timeout,
**kwargs,
)
ctx = contextvars.copy_context()
init_response = await loop.run_in_executor(None, partial(ctx.run, func))
if asyncio.iscoroutine(init_response):
return await init_response
return init_response
except Exception as e:
raise litellm.exception_type(
model="",
custom_llm_provider=custom_llm_provider or "gemini",
original_exception=e,
completion_kwargs=local_vars,
extra_kwargs=kwargs,
)
@client
def list(
custom_llm_provider: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
**kwargs,
) -> Union[AgentListResponse, Coroutine[Any, Any, AgentListResponse]]:
"""Sync: List all agents on the provider side."""
local_vars = locals()
custom_llm_provider = (
custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini"
)
try:
_is_async = kwargs.pop("alist_agents", False) is True
kwargs.setdefault("custom_llm_provider", custom_llm_provider)
litellm_params = GenericLiteLLMParams(**kwargs)
logging_obj = _make_logging_obj(
kwargs, "", custom_llm_provider, "list_agents", {}
)
config = _get_agents_api_config(custom_llm_provider)
return agents_http_handler.list_agents(
agents_api_config=config,
litellm_params=litellm_params,
logging_obj=logging_obj,
extra_headers=extra_headers,
timeout=timeout,
_is_async=_is_async,
)
except Exception as e:
raise litellm.exception_type(
model="",
custom_llm_provider=custom_llm_provider,
original_exception=e,
completion_kwargs=local_vars,
extra_kwargs=kwargs,
)
# ================================================================== #
# GET #
# ================================================================== #
@client
async def aget(
name: str,
custom_llm_provider: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
**kwargs,
) -> AgentCreateResponse:
"""Async: Get a specific agent by name."""
local_vars = locals()
try:
loop = asyncio.get_event_loop()
kwargs["aget_agent"] = True
func = partial(
get,
name=name,
custom_llm_provider=custom_llm_provider or "gemini",
extra_headers=extra_headers,
timeout=timeout,
**kwargs,
)
ctx = contextvars.copy_context()
init_response = await loop.run_in_executor(None, partial(ctx.run, func))
if asyncio.iscoroutine(init_response):
return await init_response
return init_response
except Exception as e:
raise litellm.exception_type(
model=name,
custom_llm_provider=custom_llm_provider or "gemini",
original_exception=e,
completion_kwargs=local_vars,
extra_kwargs=kwargs,
)
@client
def get(
name: str,
custom_llm_provider: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
**kwargs,
) -> Union[AgentCreateResponse, Coroutine[Any, Any, AgentCreateResponse]]:
"""Sync: Get a specific agent by name."""
local_vars = locals()
custom_llm_provider = (
custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini"
)
try:
_is_async = kwargs.pop("aget_agent", False) is True
kwargs.setdefault("custom_llm_provider", custom_llm_provider)
litellm_params = GenericLiteLLMParams(**kwargs)
logging_obj = _make_logging_obj(
kwargs, name, custom_llm_provider, "get_agent", {"name": name}
)
config = _get_agents_api_config(custom_llm_provider)
return agents_http_handler.get_agent(
agents_api_config=config,
name=name,
litellm_params=litellm_params,
logging_obj=logging_obj,
extra_headers=extra_headers,
timeout=timeout,
_is_async=_is_async,
)
except Exception as e:
raise litellm.exception_type(
model=name,
custom_llm_provider=custom_llm_provider,
original_exception=e,
completion_kwargs=local_vars,
extra_kwargs=kwargs,
)
# ================================================================== #
# DELETE #
# ================================================================== #
@client
async def adelete(
name: str,
custom_llm_provider: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
**kwargs,
) -> AgentDeleteResult:
"""Async: Delete a specific agent by name."""
local_vars = locals()
try:
loop = asyncio.get_event_loop()
kwargs["adelete_agent"] = True
func = partial(
delete,
name=name,
custom_llm_provider=custom_llm_provider or "gemini",
extra_headers=extra_headers,
timeout=timeout,
**kwargs,
)
ctx = contextvars.copy_context()
init_response = await loop.run_in_executor(None, partial(ctx.run, func))
if asyncio.iscoroutine(init_response):
return await init_response
return init_response
except Exception as e:
raise litellm.exception_type(
model=name,
custom_llm_provider=custom_llm_provider or "gemini",
original_exception=e,
completion_kwargs=local_vars,
extra_kwargs=kwargs,
)
@client
def delete(
name: str,
custom_llm_provider: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
**kwargs,
) -> Union[AgentDeleteResult, Coroutine[Any, Any, AgentDeleteResult]]:
"""Sync: Delete a specific agent by name."""
local_vars = locals()
custom_llm_provider = (
custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini"
)
try:
_is_async = kwargs.pop("adelete_agent", False) is True
kwargs.setdefault("custom_llm_provider", custom_llm_provider)
litellm_params = GenericLiteLLMParams(**kwargs)
logging_obj = _make_logging_obj(
kwargs, name, custom_llm_provider, "delete_agent", {"name": name}
)
config = _get_agents_api_config(custom_llm_provider)
return agents_http_handler.delete_agent(
agents_api_config=config,
name=name,
litellm_params=litellm_params,
logging_obj=logging_obj,
extra_headers=extra_headers,
timeout=timeout,
_is_async=_is_async,
)
except Exception as e:
raise litellm.exception_type(
model=name,
custom_llm_provider=custom_llm_provider,
original_exception=e,
completion_kwargs=local_vars,
extra_kwargs=kwargs,
)
# ================================================================== #
# LIST VERSIONS #
# ================================================================== #
@client
async def alist_versions(
name: str,
custom_llm_provider: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
**kwargs,
) -> AgentVersionsResponse:
"""Async: List versions of a specific agent."""
local_vars = locals()
try:
loop = asyncio.get_event_loop()
kwargs["alist_agent_versions"] = True
func = partial(
list_versions,
name=name,
custom_llm_provider=custom_llm_provider or "gemini",
extra_headers=extra_headers,
timeout=timeout,
**kwargs,
)
ctx = contextvars.copy_context()
init_response = await loop.run_in_executor(None, partial(ctx.run, func))
if asyncio.iscoroutine(init_response):
return await init_response
return init_response
except Exception as e:
raise litellm.exception_type(
model=name,
custom_llm_provider=custom_llm_provider or "gemini",
original_exception=e,
completion_kwargs=local_vars,
extra_kwargs=kwargs,
)
@client
def list_versions(
name: str,
custom_llm_provider: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
**kwargs,
) -> Union[AgentVersionsResponse, Coroutine[Any, Any, AgentVersionsResponse]]:
"""Sync: List versions of a specific agent."""
local_vars = locals()
custom_llm_provider = (
custom_llm_provider or kwargs.get("custom_llm_provider") or "gemini"
)
try:
_is_async = kwargs.pop("alist_agent_versions", False) is True
kwargs.setdefault("custom_llm_provider", custom_llm_provider)
litellm_params = GenericLiteLLMParams(**kwargs)
logging_obj = _make_logging_obj(
kwargs, name, custom_llm_provider, "list_agent_versions", {"name": name}
)
config = _get_agents_api_config(custom_llm_provider)
return agents_http_handler.list_agent_versions(
agents_api_config=config,
name=name,
litellm_params=litellm_params,
logging_obj=logging_obj,
extra_headers=extra_headers,
timeout=timeout,
_is_async=_is_async,
)
except Exception as e:
raise litellm.exception_type(
model=name,
custom_llm_provider=custom_llm_provider,
original_exception=e,
completion_kwargs=local_vars,
extra_kwargs=kwargs,
)

View file

@ -0,0 +1,23 @@
"""
Utility functions for the Agents API SDK.
"""
from typing import Optional
from litellm.llms.base_llm.agents.transformation import BaseAgentsAPIConfig
def get_provider_agents_api_config(
custom_llm_provider: Optional[str],
) -> Optional[BaseAgentsAPIConfig]:
"""
Return a provider-specific BaseAgentsAPIConfig if the provider has a
native agent-creation API, or None otherwise.
"""
from litellm.types.utils import LlmProviders
if custom_llm_provider == LlmProviders.GEMINI.value:
from litellm.llms.gemini.agents.transformation import GeminiAgentsConfig
return GeminiAgentsConfig()
return None

View file

@ -41,27 +41,55 @@ from litellm.types.interactions import (
from litellm.types.router import GenericLiteLLMParams
class InteractionsHTTPHandler:
class _BaseHTTPHandler:
"""
Shared HTTP infrastructure for LiteLLM handler classes.
Provides common client resolution and error-mapping helpers so that
handler subclasses (InteractionsHTTPHandler, AgentsHTTPHandler, …) do
not duplicate this boilerplate.
"""
def _handle_error(self, e: Exception, provider_config: Any) -> Exception:
if isinstance(e, httpx.HTTPStatusError):
return provider_config.get_error_class(
error_message=e.response.text,
status_code=e.response.status_code,
headers=dict(e.response.headers),
)
return e
def _sync_client(
self,
litellm_params: GenericLiteLLMParams,
client: Optional[HTTPHandler],
) -> HTTPHandler:
return client or _get_httpx_client(
params={"ssl_verify": litellm_params.get("ssl_verify", None)}
)
def _async_client(
self,
litellm_params: GenericLiteLLMParams,
client: Optional[AsyncHTTPHandler],
) -> AsyncHTTPHandler:
# GenericLiteLLMParams.get uses getattr; an unset field is None, not the default.
custom_llm_provider = litellm_params.get("custom_llm_provider") or "gemini"
return client or get_async_httpx_client(
llm_provider=litellm.LlmProviders(custom_llm_provider),
params={"ssl_verify": litellm_params.get("ssl_verify", None)},
)
class InteractionsHTTPHandler(_BaseHTTPHandler):
"""
HTTP handler for Interactions API requests.
"""
def _handle_error(
self,
e: Exception,
provider_config: BaseInteractionsAPIConfig,
) -> Exception:
"""Handle errors from HTTP requests."""
if isinstance(e, httpx.HTTPStatusError):
error_message = e.response.text
status_code = e.response.status_code
headers = dict(e.response.headers)
return provider_config.get_error_class(
error_message=error_message,
status_code=status_code,
headers=headers,
)
return e
# _handle_error is inherited from _BaseHTTPHandler (accepts Any provider_config).
# AgentsHTTPHandler also extends this class and passes BaseAgentsAPIConfig, which
# is structurally compatible but a different type — keeping the override here with
# BaseInteractionsAPIConfig would cause type errors in the subclass.
# =========================================================
# CREATE INTERACTION

View file

@ -2,7 +2,17 @@
Streaming iterator for transforming Responses API stream to Interactions API stream.
"""
from typing import Any, AsyncIterator, Dict, Iterator, Optional, cast
from collections import deque
from typing import (
Any,
AsyncIterator,
Deque,
Dict,
Iterator,
List,
Optional,
cast,
)
from litellm.responses.streaming_iterator import (
BaseResponsesAPIStreamingIterator,
@ -29,7 +39,13 @@ class LiteLLMResponsesInteractionsStreamingIterator:
This class handles both sync and async iteration, transforming Responses API
streaming events (output.text.delta, response.completed, etc.) to Interactions
API streaming events (content.delta, interaction.complete, etc.).
API streaming events.
Schema selection:
- New schema (default, use_legacy_interactions_schema=False):
interaction.created -> step.start -> step.delta ... -> step.stop -> interaction.completed
- Legacy schema (use_legacy_interactions_schema=True, remove after June 8 2026):
interaction.start -> content.start -> content.delta ... -> content.stop -> interaction.complete
"""
def __init__(
@ -41,6 +57,8 @@ class LiteLLMResponsesInteractionsStreamingIterator:
custom_llm_provider: Optional[str] = None,
litellm_metadata: Optional[Dict[str, Any]] = None,
):
import litellm
self.model = model
self.responses_stream_iterator = litellm_custom_stream_wrapper
self.request_input = request_input
@ -51,66 +69,156 @@ class LiteLLMResponsesInteractionsStreamingIterator:
self.collected_text = ""
self.sent_interaction_start = False
self.sent_content_start = False
# Capture the schema flag once at construction time so all events
# emitted by this stream use a consistent schema, even if the global
# flag is mutated mid-stream (e.g. by a config reload).
self._use_legacy: bool = litellm.use_legacy_interactions_schema
# Buffer of events that have been derived from upstream chunks but not
# yet returned to the caller. A single Responses API chunk may expand
# into multiple Interactions API events (e.g. the first text delta
# produces interaction.created + step.start + step.delta), and the
# terminal sequence on stream end may also span multiple events
# (step.stop + interaction.completed).
self._pending_events: Deque[InteractionsAPIStreamingResponse] = deque()
# Tracks whether we've already emitted a terminal completion event so
# the StopIteration fallback path doesn't double-emit.
self._sent_completion_event = False
# ID resolved from the first upstream chunk (item_id on a text delta or
# response.id on response.created). Persisted so the EOF terminal
# events stay correlated with the start events delivered earlier.
self._interaction_id: Optional[str] = None
def _transform_responses_chunk_to_interactions_chunk(
self,
responses_chunk: ResponsesAPIStreamingResponse,
) -> Optional[InteractionsAPIStreamingResponse]:
# ------------------------------------------------------------------
# Event builders
# ------------------------------------------------------------------
def _build_interaction_start_event(
self, interaction_id: str
) -> InteractionsAPIStreamingResponse:
event_type = "interaction.start" if self._use_legacy else "interaction.created"
return InteractionsAPIStreamingResponse(
event_type=event_type,
id=interaction_id,
object="interaction",
status="in_progress",
model=self.model,
)
def _build_content_start_event(
self, interaction_id: str
) -> InteractionsAPIStreamingResponse:
if self._use_legacy:
return InteractionsAPIStreamingResponse(
event_type="content.start",
id=interaction_id,
object="content",
delta={"type": "text", "text": ""},
)
return InteractionsAPIStreamingResponse(
event_type="step.start",
index=0,
step={"type": "model_output", "content": []},
)
def _build_text_delta_event(
self, interaction_id: str, delta_text: str
) -> InteractionsAPIStreamingResponse:
if self._use_legacy:
return InteractionsAPIStreamingResponse(
event_type="content.delta",
id=interaction_id,
object="content",
delta={"type": "text", "text": delta_text},
)
return InteractionsAPIStreamingResponse(
event_type="step.delta",
index=0,
delta={"type": "text", "text": delta_text},
)
def _build_content_stop_event(
self, interaction_id: Optional[str]
) -> InteractionsAPIStreamingResponse:
if self._use_legacy:
return InteractionsAPIStreamingResponse(
event_type="content.stop",
id=interaction_id,
object="content",
delta={"type": "text", "text": self.collected_text},
)
return InteractionsAPIStreamingResponse(
event_type="step.stop",
index=0,
)
def _build_completion_event(
self, response_id: str
) -> InteractionsAPIStreamingResponse:
if self._use_legacy:
return InteractionsAPIStreamingResponse(
event_type="interaction.complete",
id=response_id,
object="interaction",
status="completed",
model=self.model,
outputs=[{"type": "text", "text": self.collected_text}],
)
return InteractionsAPIStreamingResponse(
event_type="interaction.completed",
id=response_id,
object="interaction",
status="completed",
model=self.model,
steps=[
{
"type": "model_output",
"content": [{"type": "text", "text": self.collected_text}],
}
],
)
# ------------------------------------------------------------------
# Per-chunk transform (returns a list of events to enqueue)
# ------------------------------------------------------------------
def _events_for_chunk(
self, responses_chunk: ResponsesAPIStreamingResponse
) -> List[InteractionsAPIStreamingResponse]:
"""
Transform a Responses API streaming chunk to an Interactions API streaming chunk.
Translate a single upstream Responses API chunk into the list of
Interactions API events it should produce.
Responses API events:
- output.text.delta -> content.delta
- response.completed -> interaction.complete
Interactions API events:
- interaction.start
- content.start
- content.delta
- content.stop
- interaction.complete
Returning a list (rather than a single event) lets a chunk emit any
synthetic start events that haven't been sent yet *together with* the
actual delta event, so we never silently drop the chunk's payload.
"""
if not responses_chunk:
return None
return []
# Handle OutputTextDeltaEvent -> content.delta
# Text delta: emit any missing start events, then the delta itself.
if isinstance(responses_chunk, OutputTextDeltaEvent):
delta_text = (
responses_chunk.delta if isinstance(responses_chunk.delta, str) else ""
)
self.collected_text += delta_text
interaction_id = (
getattr(responses_chunk, "item_id", None) or f"interaction_{id(self)}"
)
if self._interaction_id is None:
self._interaction_id = interaction_id
# Send interaction.start if not sent
events: List[InteractionsAPIStreamingResponse] = []
if not self.sent_interaction_start:
self.sent_interaction_start = True
return InteractionsAPIStreamingResponse(
event_type="interaction.start",
id=getattr(responses_chunk, "item_id", None)
or f"interaction_{id(self)}",
object="interaction",
status="in_progress",
model=self.model,
)
# Send content.start if not sent
events.append(self._build_interaction_start_event(interaction_id))
if not self.sent_content_start:
self.sent_content_start = True
return InteractionsAPIStreamingResponse(
event_type="content.start",
id=getattr(responses_chunk, "item_id", None),
object="content",
delta={"type": "text", "text": ""},
)
events.append(self._build_content_start_event(interaction_id))
events.append(self._build_text_delta_event(interaction_id, delta_text))
return events
# Send content.delta
return InteractionsAPIStreamingResponse(
event_type="content.delta",
id=getattr(responses_chunk, "item_id", None),
object="content",
delta={"text": delta_text},
)
# Handle ResponseCreatedEvent or ResponseInProgressEvent -> interaction.start
# Response created / in-progress: synthesize interaction start if we
# haven't already sent one.
if isinstance(responses_chunk, (ResponseCreatedEvent, ResponseInProgressEvent)):
if not self.sent_interaction_start:
self.sent_interaction_start = True
@ -118,169 +226,136 @@ class LiteLLMResponsesInteractionsStreamingIterator:
getattr(responses_chunk.response, "id", None)
if hasattr(responses_chunk, "response")
else None
)
return InteractionsAPIStreamingResponse(
event_type="interaction.start",
id=response_id or f"interaction_{id(self)}",
object="interaction",
status="in_progress",
model=self.model,
)
) or f"interaction_{id(self)}"
if self._interaction_id is None:
self._interaction_id = response_id
return [self._build_interaction_start_event(response_id)]
return []
# Handle ResponseCompletedEvent -> interaction.complete
# Response completed: emit step.stop (if content was started) followed
# by the terminal completion event. Prefer the interaction id already
# established by earlier events so consumers can correlate the start
# and completion events by id (response.id may differ from the item_id
# used to derive the initial id when the stream starts directly with a
# text delta).
if isinstance(responses_chunk, ResponseCompletedEvent):
self.finished = True
response = responses_chunk.response
# Send content.stop first if content was started
if self.sent_content_start:
# Note: We'll send this in the iterator, not here
pass
# Send interaction.complete
return InteractionsAPIStreamingResponse(
event_type="interaction.complete",
id=getattr(response, "id", None) or f"interaction_{id(self)}",
object="interaction",
status="completed",
model=self.model,
outputs=[
{
"type": "text",
"text": self.collected_text,
}
],
response_id = (
self._interaction_id
or getattr(response, "id", None)
or f"interaction_{id(self)}"
)
# For other event types, return None (skip)
return None
terminal: List[InteractionsAPIStreamingResponse] = []
if self.sent_content_start:
terminal.append(self._build_content_stop_event(response_id))
terminal.append(self._build_completion_event(response_id))
self._sent_completion_event = True
return terminal
return []
def _build_terminal_events_on_eof(
self,
) -> List[InteractionsAPIStreamingResponse]:
"""
Build the events to flush when the upstream stream ends without a
ResponseCompletedEvent. Ensures consumers always observe a terminal
interaction.completed/interaction.complete carrying the full text.
"""
if self._sent_completion_event:
return []
fallback_id = self._interaction_id or f"interaction_{id(self)}"
terminal: List[InteractionsAPIStreamingResponse] = []
if self.sent_content_start:
terminal.append(self._build_content_stop_event(fallback_id))
if self.sent_interaction_start or self.collected_text:
terminal.append(self._build_completion_event(fallback_id))
self._sent_completion_event = True
return terminal
# ------------------------------------------------------------------
# Iteration
# ------------------------------------------------------------------
def __iter__(self) -> Iterator[InteractionsAPIStreamingResponse]:
"""Sync iterator implementation."""
return self
def __next__(self) -> InteractionsAPIStreamingResponse:
"""Get next chunk in sync mode."""
if self._pending_events:
return self._pending_events.popleft()
if self.finished:
raise StopIteration
# Check if we have a pending interaction.complete to send
if hasattr(self, "_pending_interaction_complete"):
pending: InteractionsAPIStreamingResponse = getattr(
self, "_pending_interaction_complete"
)
delattr(self, "_pending_interaction_complete")
return pending
# Use a loop instead of recursion to avoid stack overflow
sync_iterator = cast(
SyncResponsesAPIStreamingIterator, self.responses_stream_iterator
)
while True:
try:
# Get next chunk from responses API stream
chunk = next(sync_iterator)
# Transform chunk (chunk is already a ResponsesAPIStreamingResponse)
transformed = self._transform_responses_chunk_to_interactions_chunk(
chunk
)
if transformed:
# If we finished and content was started, send content.stop before interaction.complete
if (
self.finished
and self.sent_content_start
and transformed.event_type == "interaction.complete"
):
# Send content.stop first
content_stop = InteractionsAPIStreamingResponse(
event_type="content.stop",
id=transformed.id,
object="content",
delta={"type": "text", "text": self.collected_text},
)
# Store the interaction.complete to send next
self._pending_interaction_complete = transformed
return content_stop
return transformed
# If no transformation, continue to next chunk (loop continues)
except StopIteration:
self.finished = True
self._pending_events.extend(self._build_terminal_events_on_eof())
if self._pending_events:
return self._pending_events.popleft()
raise
# Send final events if needed
if self.sent_content_start:
return InteractionsAPIStreamingResponse(
event_type="content.stop",
object="content",
delta={"type": "text", "text": self.collected_text},
)
raise StopIteration
events = self._events_for_chunk(chunk)
if events:
self._pending_events.extend(events)
return self._pending_events.popleft()
def __aiter__(self) -> AsyncIterator[InteractionsAPIStreamingResponse]:
"""Async iterator implementation."""
return self
async def __anext__(self) -> InteractionsAPIStreamingResponse:
"""Get next chunk in async mode."""
if self._pending_events:
return self._pending_events.popleft()
if self.finished:
raise StopAsyncIteration
# Check if we have a pending interaction.complete to send
if hasattr(self, "_pending_interaction_complete"):
pending: InteractionsAPIStreamingResponse = getattr(
self, "_pending_interaction_complete"
)
delattr(self, "_pending_interaction_complete")
return pending
# Use a loop instead of recursion to avoid stack overflow
async_iterator = cast(
ResponsesAPIStreamingIterator, self.responses_stream_iterator
)
while True:
try:
# Get next chunk from responses API stream
chunk = await async_iterator.__anext__()
# Transform chunk (chunk is already a ResponsesAPIStreamingResponse)
transformed = self._transform_responses_chunk_to_interactions_chunk(
chunk
)
if transformed:
# If we finished and content was started, send content.stop before interaction.complete
if (
self.finished
and self.sent_content_start
and transformed.event_type == "interaction.complete"
):
# Send content.stop first
content_stop = InteractionsAPIStreamingResponse(
event_type="content.stop",
id=transformed.id,
object="content",
delta={"type": "text", "text": self.collected_text},
)
# Store the interaction.complete to send next
self._pending_interaction_complete = transformed
return content_stop
return transformed
# If no transformation, continue to next chunk (loop continues)
except StopAsyncIteration:
self.finished = True
self._pending_events.extend(self._build_terminal_events_on_eof())
if self._pending_events:
return self._pending_events.popleft()
raise
# Send final events if needed
if self.sent_content_start:
return InteractionsAPIStreamingResponse(
event_type="content.stop",
object="content",
delta={"type": "text", "text": self.collected_text},
)
events = self._events_for_chunk(chunk)
if events:
self._pending_events.extend(events)
return self._pending_events.popleft()
raise StopAsyncIteration
# ------------------------------------------------------------------
# Backwards-compatible single-chunk transform (used by tests and any
# external callers that drove the iterator chunk-by-chunk pre-fix).
# ------------------------------------------------------------------
def _transform_responses_chunk_to_interactions_chunk(
self,
responses_chunk: ResponsesAPIStreamingResponse,
) -> Optional[InteractionsAPIStreamingResponse]:
"""
Compatibility shim: returns the *first* event produced for this chunk
and queues any remaining events on ``self._pending_events`` so they
are surfaced on subsequent calls/iterations.
Prefer ``_events_for_chunk`` in new code.
"""
events = self._events_for_chunk(responses_chunk)
if not events:
return None
first = events[0]
if len(events) > 1:
self._pending_events.extend(events[1:])
return first

View file

@ -226,29 +226,37 @@ class LiteLLMResponsesInteractionsConfig:
- Map status
- Extract usage
"""
# Extract text from outputs
outputs = []
# Extract text from outputs and build both `outputs` (legacy) and `steps` (new schema).
outputs: List[Dict[str, Any]] = []
steps: List[Dict[str, Any]] = []
if hasattr(responses_response, "output") and responses_response.output:
for output_item in responses_response.output:
# Use getattr with None default to safely access content
content = getattr(output_item, "content", None)
if content is not None:
content_items = content if isinstance(content, list) else [content]
model_output_contents: List[Dict[str, Any]] = []
for content_item in content_items:
# Check if content_item has text attribute
text = getattr(content_item, "text", None)
if text is not None:
outputs.append(
{
"type": "text",
"text": text,
}
)
# Use independent dict instances so mutations to one
# of `outputs` / `steps` don't leak into the other.
outputs.append({"type": "text", "text": text})
model_output_contents.append({"type": "text", "text": text})
elif (
isinstance(content_item, dict)
and content_item.get("type") == "text"
):
outputs.append(content_item)
outputs.append({**content_item})
model_output_contents.append({**content_item})
if model_output_contents:
steps.append(
{
"type": "model_output",
"content": model_output_contents,
}
)
# Convert created_at to ISO string
created_at = getattr(responses_response, "created_at", None)
@ -270,12 +278,14 @@ class LiteLLMResponsesInteractionsConfig:
else:
interactions_status = status
# Build interactions response
# Build interactions response — populate both `outputs` (legacy schema) and
# `steps` (new schema) so callers work regardless of which schema they expect.
interactions_response_dict: Dict[str, Any] = {
"id": getattr(responses_response, "id", ""),
"object": "interaction",
"status": interactions_status,
"outputs": outputs,
"steps": steps,
"model": model or getattr(responses_response, "model", ""),
"created": created,
}

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="...")
"""
@ -48,6 +48,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.types.interactions import (
CancelInteractionResult,
DeleteInteractionResult,
InteractionEnvironment,
InteractionInput,
InteractionsAPIResponse,
InteractionsAPIStreamingResponse,
@ -80,6 +81,8 @@ async def acreate(
store: Optional[bool] = None,
# Background execution
background: Optional[bool] = None,
# Agent execution environment ("remote", env id, or remote config object)
environment: Optional[InteractionEnvironment] = None,
# Response format
response_modalities: Optional[List[str]] = None,
response_format: Optional[Dict[str, Any]] = None,
@ -109,6 +112,10 @@ async def acreate(
stream: Whether to stream the response
store: Whether to store the response for later retrieval
background: Whether to run in background
environment: Agent execution environment — ``"remote"``, an existing env id
string, or a config object such as
``{"type": "remote", "sources": [...]}`` /
``{"type": "remote", "network": {...}}``
response_modalities: Requested response modalities (TEXT, IMAGE, AUDIO)
response_format: JSON schema for response format
response_mime_type: MIME type of the response
@ -144,6 +151,7 @@ async def acreate(
stream=stream,
store=store,
background=background,
environment=environment,
response_modalities=response_modalities,
response_format=response_format,
response_mime_type=response_mime_type,
@ -194,6 +202,8 @@ def create(
store: Optional[bool] = None,
# Background execution
background: Optional[bool] = None,
# Agent execution environment ("remote", env id, or remote config object)
environment: Optional[InteractionEnvironment] = None,
# Response format
response_modalities: Optional[List[str]] = None,
response_format: Optional[Dict[str, Any]] = None,
@ -231,6 +241,10 @@ def create(
stream: Whether to stream the response
store: Whether to store the response for later retrieval
background: Whether to run in background
environment: Agent execution environment — ``"remote"``, an existing env id
string, or a config object such as
``{"type": "remote", "sources": [...]}`` /
``{"type": "remote", "network": {...}}``
response_modalities: Requested response modalities (TEXT, IMAGE, AUDIO)
response_format: JSON schema for response format
response_mime_type: MIME type of the response
@ -252,7 +266,14 @@ def create(
litellm_params = GenericLiteLLMParams(**kwargs)
if model:
# Routing logic:
# - agent provided (no model, or model accidentally set to agent name) → gemini
# - model provided → resolve provider via get_llm_provider (normal routing)
if agent and model == agent:
model = None
if agent and not model:
custom_llm_provider = custom_llm_provider or "gemini"
elif model:
model, custom_llm_provider, _, _ = litellm.get_llm_provider(
model=model,
custom_llm_provider=custom_llm_provider,

View file

@ -101,10 +101,14 @@ class BaseInteractionsAPIStreamingIterator:
)
)
# Store the completed response (check for status=completed)
if (
streaming_response
and getattr(streaming_response, "status", None) == "completed"
# Store the completed response.
# Legacy schema signals completion via status="completed".
# New schema (Api-Revision: 2026-05-20) uses event_type="interaction.completed".
# Remove the legacy check after June 8, 2026.
if streaming_response and (
getattr(streaming_response, "status", None) == "completed"
or getattr(streaming_response, "event_type", None)
== "interaction.completed"
):
self.completed_response = streaming_response
self._handle_logging_completed_response()

View file

@ -15,6 +15,7 @@ INTERACTIONS_API_OPTIONAL_PARAMS = {
"stream",
"store",
"background",
"environment",
"response_modalities",
"response_format",
"response_mime_type",

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(
@ -1769,9 +1763,12 @@ class Logging(LiteLLMLoggingBaseClass):
self.model_call_details["response_cost"] = 0.0
elif "response_cost" in hidden_params:
self.model_call_details["response_cost"] = hidden_params["response_cost"]
elif self.model_call_details.get("response_cost") is not None:
elif (
existing_cost := self.model_call_details.get("response_cost")
) is not None and existing_cost != 0:
# Preserve response_cost if already calculated (e.g., by pass-through
# handlers like Gemini/Vertex which call completion_cost directly)
# handlers like Gemini/Vertex which call completion_cost directly).
# Do not preserve 0 from failure_handler on intermediate router retries.
pass
else:
self.model_call_details["response_cost"] = self._response_cost_calculator(

View file

@ -1233,6 +1233,7 @@ def infer_protocol_value(
def _gemini_tool_call_invoke_helper(
function_call_params: ChatCompletionToolCallFunctionChunk,
tool_call_id: Optional[str] = None,
) -> Optional[VertexFunctionCall]:
name = function_call_params.get("name", "") or ""
arguments = function_call_params.get("arguments", "")
@ -1248,6 +1249,10 @@ def _gemini_tool_call_invoke_helper(
name=name,
args=arguments_dict,
)
if tool_call_id:
clean_id = tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1)[0]
if clean_id:
function_call["id"] = clean_id
return function_call
@ -1339,6 +1344,7 @@ def _get_dummy_thought_signature() -> str:
def convert_to_gemini_tool_call_invoke(
message: ChatCompletionAssistantMessage,
model: Optional[str] = None,
custom_llm_provider: Optional[str] = None,
) -> List[VertexPartType]:
"""
OpenAI tool invokes:
@ -1384,12 +1390,26 @@ def convert_to_gemini_tool_call_invoke(
tool_calls = message.get("tool_calls", None)
function_call = message.get("function_call", None)
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
forward_tool_call_id = bool(
model
and VertexGeminiConfig._forward_gemini_function_call_id(
model, custom_llm_provider
)
)
if tool_calls is not None:
for idx, tool in enumerate(tool_calls):
if "function" in tool:
gemini_function_call: Optional[VertexFunctionCall] = (
_gemini_tool_call_invoke_helper(
function_call_params=tool["function"]
function_call_params=tool["function"],
tool_call_id=(
tool.get("id") if forward_tool_call_id else None
),
)
)
if gemini_function_call is not None:
@ -1429,10 +1449,6 @@ def convert_to_gemini_tool_call_invoke(
thought_signature = provider_fields.get("thought_signature")
# If no signature found and model is gemini-3, use dummy signature
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
if (
not thought_signature
and model
@ -1462,6 +1478,8 @@ def convert_to_gemini_tool_call_invoke(
def convert_to_gemini_tool_call_result( # noqa: PLR0915
message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage],
last_message_with_tool_calls: Optional[dict],
model: Optional[str] = None,
custom_llm_provider: Optional[str] = None,
) -> Union[VertexPartType, List[VertexPartType]]:
"""
OpenAI message with a tool result looks like:
@ -1602,6 +1620,23 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915
):
name = tool.get("function", {}).get("name", "")
# Echo the OpenAI tool_call_id on functionResponse (strip thought-signature suffix).
# Only Google AI Studio Gemini 3+ accepts `id` on function_response parts.
# Vertex AI and older Gemini models reject the field with HTTP 400.
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
gemini_call_id: Optional[str] = None
if model and VertexGeminiConfig._forward_gemini_function_call_id(
model, custom_llm_provider
):
raw_tool_call_id = message.get("tool_call_id")
if raw_tool_call_id and isinstance(raw_tool_call_id, str):
stripped_id = raw_tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1)[0]
if stripped_id:
gemini_call_id = stripped_id
if not name:
raise Exception(
"Missing corresponding tool call for tool response message. Received - message={}, last_message_with_tool_calls={}".format(
@ -1632,6 +1667,8 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915
name=name,
response=response_data, # type: ignore
)
if gemini_call_id:
_function_response["id"] = gemini_call_id
# Create part with function_response, and optionally inline_data for images (Computer Use)
_part: VertexPartType = {"function_response": _function_response}
@ -5553,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

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

@ -1506,9 +1506,21 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
optional_params["metadata"] = {"user_id": value}
elif param == "thinking":
optional_params["thinking"] = value
elif param == "reasoning_effort" and isinstance(value, str):
elif param == "reasoning_effort":
# Accept both string ("low") and dict ({"effort": "low",
# "summary": "concise"}). The Responses->Chat parser keeps the
# full dict when `summary` is set (see #25359), so a dict here
# is the standard shape Otto/OpenAI-Responses-Bridge callers
# send. Coerce to the effort string before mapping — same
# shape-tolerance the GPT-5 path already implements in
# `_normalize_reasoning_effort_for_chat_completion`.
effort_value = value
if isinstance(effort_value, dict):
effort_value = effort_value.get("effort")
if not isinstance(effort_value, str):
continue
mapped_thinking = AnthropicConfig._map_reasoning_effort(
reasoning_effort=value,
reasoning_effort=effort_value,
model=model,
llm_provider=self.custom_llm_provider or "anthropic",
)
@ -1519,12 +1531,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
optional_params["thinking"] = mapped_thinking
if AnthropicConfig._is_adaptive_thinking_model(model):
mapped_effort = REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT.get(
value
effort_value
)
if mapped_effort is None:
AnthropicConfig._raise_invalid_reasoning_effort(
model=model,
value=value,
value=effort_value,
llm_provider=self.custom_llm_provider or "anthropic",
)
optional_params["output_config"] = {"effort": mapped_effort}

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

@ -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,9 +1,16 @@
from typing import Optional
from urllib.parse import parse_qs, urlparse, urlunparse
from litellm.llms.azure.common_utils import BaseAzureLLM
from litellm.llms.openai.containers.transformation import OpenAIContainerConfig
from litellm.types.router import GenericLiteLLMParams
# Endpoint-specific path suffixes that may appear in a deployment's api_base
# (e.g. the responses endpoint URL is stored as api_base for Azure models).
# Strip these before building the containers URL so we always start from the
# resource root (https://resource.cognitiveservices.azure.com).
_AZURE_ENDPOINT_PATHS = ("/openai/responses",)
class AzureContainerConfig(OpenAIContainerConfig):
"""
@ -27,6 +34,27 @@ class AzureContainerConfig(OpenAIContainerConfig):
litellm_params=GenericLiteLLMParams(api_key=api_key),
)
@staticmethod
def _normalize_api_base(api_base: Optional[str]) -> Optional[str]:
"""Strip endpoint-specific path suffixes from api_base to get the resource root."""
if not api_base:
return api_base
parsed = urlparse(api_base)
path = parsed.path.rstrip("/")
for ep in _AZURE_ENDPOINT_PATHS:
if path.endswith(ep):
return urlunparse(
(parsed.scheme, parsed.netloc, path[: -len(ep)], "", "", "")
)
return api_base
@staticmethod
def _extract_api_version(api_base: Optional[str]) -> Optional[str]:
"""Return the api-version query param from api_base if present."""
if not api_base:
return None
return parse_qs(urlparse(api_base).query).get("api-version", [None])[0]
def get_complete_url(
self,
api_base: Optional[str],
@ -39,10 +67,19 @@ class AzureContainerConfig(OpenAIContainerConfig):
{endpoint}/openai/v1/containers
when api_version is 'v1', 'latest', or 'preview'; otherwise:
{endpoint}/openai/containers
The deployment's api_base may be the responses endpoint URL
(e.g. .../openai/responses?api-version=2025-04-01-preview). We
prefer the api-version embedded there over the deployment's
api_version field, which may point to an older chat API version.
"""
effective_params = dict(litellm_params)
api_version_from_base = self._extract_api_version(api_base)
if api_version_from_base:
effective_params["api_version"] = api_version_from_base
return BaseAzureLLM._get_base_azure_url(
api_base=api_base,
litellm_params=litellm_params,
api_base=self._normalize_api_base(api_base),
litellm_params=effective_params,
route="/openai/containers",
default_api_version="v1",
)

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

View file

@ -0,0 +1,165 @@
"""
Base transformation class for provider-side Agents API.
Providers that have a native agents CRUD API (e.g. Gemini v1beta/agents)
subclass BaseAgentsAPIConfig and implement the abstract methods.
The HTTP calls are handled by AgentsHTTPHandler — this class is pure
transform logic (same separation as BaseInteractionsAPIConfig /
InteractionsHTTPHandler).
"""
from abc import ABC, abstractmethod
from typing import Any, Dict, Optional, Tuple, Union
import httpx
from litellm.types.agents import (
AgentCreateResponse,
AgentDeleteResult,
AgentListResponse,
AgentVersionsResponse,
)
class BaseAgentsAPIConfig(ABC):
"""
Minimal interface for providers that expose a native agents CRUD API.
"""
# ------------------------------------------------------------------ #
# CREATE #
# ------------------------------------------------------------------ #
@abstractmethod
def get_complete_url(
self,
api_base: Optional[str],
litellm_params: Dict[str, Any],
) -> str:
"""Return the full URL for POST /agents (create)."""
@abstractmethod
def validate_environment(
self,
headers: Dict[str, str],
litellm_params: Dict[str, Any],
) -> Dict[str, str]:
"""Validate credentials and return auth headers."""
@abstractmethod
def transform_create_request(
self,
name: str,
litellm_params: Dict[str, Any],
) -> Dict[str, Any]:
"""Map name + litellm_params to the provider's create-agent body."""
@abstractmethod
def transform_create_response(
self,
raw_response: httpx.Response,
name: str,
) -> AgentCreateResponse:
"""Parse create response. Raise on non-2xx."""
# ------------------------------------------------------------------ #
# LIST #
# ------------------------------------------------------------------ #
@abstractmethod
def transform_list_request(
self,
api_base: Optional[str],
litellm_params: Dict[str, Any],
) -> Tuple[str, Dict[str, Any]]:
"""Return (url, query_params) for GET /agents."""
@abstractmethod
def transform_list_response(
self,
raw_response: httpx.Response,
) -> AgentListResponse:
"""Parse list-agents response. Raise on non-2xx."""
# ------------------------------------------------------------------ #
# GET #
# ------------------------------------------------------------------ #
@abstractmethod
def transform_get_request(
self,
name: str,
api_base: Optional[str],
litellm_params: Dict[str, Any],
) -> Tuple[str, Dict[str, Any]]:
"""Return (url, query_params) for GET /agents/{name}."""
@abstractmethod
def transform_get_response(
self,
raw_response: httpx.Response,
name: str,
) -> AgentCreateResponse:
"""Parse get-agent response. Raise on non-2xx."""
# ------------------------------------------------------------------ #
# DELETE #
# ------------------------------------------------------------------ #
@abstractmethod
def transform_delete_request(
self,
name: str,
api_base: Optional[str],
litellm_params: Dict[str, Any],
) -> str:
"""Return the URL for DELETE /agents/{name}."""
@abstractmethod
def transform_delete_response(
self,
raw_response: httpx.Response,
name: str,
) -> AgentDeleteResult:
"""Parse delete-agent response. Raise on non-2xx."""
# ------------------------------------------------------------------ #
# LIST VERSIONS #
# ------------------------------------------------------------------ #
@abstractmethod
def transform_list_versions_request(
self,
name: str,
api_base: Optional[str],
litellm_params: Dict[str, Any],
) -> Tuple[str, Dict[str, Any]]:
"""Return (url, query_params) for GET /agents/{name}/versions."""
@abstractmethod
def transform_list_versions_response(
self,
raw_response: httpx.Response,
name: str,
) -> AgentVersionsResponse:
"""Parse list-versions response. Raise on non-2xx."""
# ------------------------------------------------------------------ #
# ERROR HANDLING #
# ------------------------------------------------------------------ #
def get_error_class(
self,
error_message: str,
status_code: int,
headers: Union[dict, httpx.Headers],
) -> Exception:
"""Map HTTP error status codes to provider-specific exceptions."""
from litellm.llms.base_llm.chat.transformation import BaseLLMException
return BaseLLMException(
status_code=status_code,
message=error_message,
headers=headers,
)

View file

@ -54,6 +54,7 @@ class OCRUsageInfo(LiteLLMPydanticObjectBase):
"""Usage information from OCR response."""
pages_processed: Optional[int] = None
credits: Optional[float] = None
doc_size_bytes: Optional[int] = None
model_config = {"extra": "allow"}

View file

@ -299,9 +299,9 @@ class AmazonInvokeAgentConfig(BaseConfig, BaseAWSLLM):
)
def _get_response_stream_shape(self):
from litellm.llms.bedrock.common_utils import BEDROCK_RESPONSE_STREAM_SHAPE
from litellm.llms.bedrock.common_utils import get_bedrock_response_stream_shape
return BEDROCK_RESPONSE_STREAM_SHAPE
return get_bedrock_response_stream_shape()
def _extract_response_content(self, events: InvokeAgentEventList) -> str:
"""Extract the final response content from parsed events."""

View file

@ -68,9 +68,9 @@ from litellm.utils import CustomStreamWrapper, get_secret
from ..base_aws_llm import BaseAWSLLM
from ..common_utils import (
BEDROCK_RESPONSE_STREAM_SHAPE,
BedrockError,
ModelResponseIterator,
get_bedrock_response_stream_shape,
get_bedrock_tool_name,
)
@ -1828,7 +1828,8 @@ class AWSEventStreamDecoder:
yield self._chunk_parser(chunk_data=_data)
def _parse_message_from_event(self, event) -> Optional[str]:
if BEDROCK_RESPONSE_STREAM_SHAPE is None:
response_stream_shape = get_bedrock_response_stream_shape()
if response_stream_shape is None:
raise BedrockError(
status_code=500,
message=(
@ -1837,9 +1838,7 @@ class AWSEventStreamDecoder:
),
)
response_dict = event.to_response_dict()
parsed_response = self.parser.parse(
response_dict, BEDROCK_RESPONSE_STREAM_SHAPE
)
parsed_response = self.parser.parse(response_dict, response_stream_shape)
if response_dict["status_code"] != 200:
decoded_body = response_dict["body"].decode()

View file

@ -3,6 +3,7 @@ from typing import TYPE_CHECKING, Any, List, Optional
import httpx
from litellm.anthropic_beta_headers_manager import filter_and_transform_beta_headers
from litellm.litellm_core_utils.litellm_logging import verbose_logger
from litellm.litellm_core_utils.prompt_templates.factory import (
convert_to_anthropic_image_obj,
)
@ -22,6 +23,7 @@ from litellm.llms.bedrock.common_utils import (
from litellm.types.llms.anthropic import ANTHROPIC_TOOL_SEARCH_BETA_HEADER
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponse
from litellm.utils import _supports_factory
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
@ -169,6 +171,24 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
anthropic_request.pop("model", None)
anthropic_request.pop("stream", None)
anthropic_request.pop("output_format", None)
if not (
_supports_factory(
model=model,
custom_llm_provider="bedrock",
key="supports_output_config",
)
or AnthropicConfig._model_supports_effort_param(model)
):
if anthropic_request.pop("output_config", None) is not None:
verbose_logger.warning(
"Bedrock Invoke: stripping unsupported `output_config` for "
"model=%s — neither `supports_output_config` nor any "
"`supports_*_reasoning_effort` flag is set in "
"model_prices_and_context_window.json. Add the capability "
"flag to the model JSON entry if this model accepts "
"`output_config`.",
model,
)
if "anthropic_version" not in anthropic_request:
anthropic_request["anthropic_version"] = self.anthropic_version

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

@ -4,6 +4,7 @@ from __future__ import annotations
Common utilities used across bedrock chat/embedding/image generation
"""
import functools
import json
import os
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union
@ -963,10 +964,8 @@ def _load_bedrock_response_stream_shape():
"""
Load the ResponseStream shape from botocore's bundled bedrock-runtime schema.
Called once at module import time; the result is stored in
``BEDROCK_RESPONSE_STREAM_SHAPE`` and reused for the process lifetime.
Returns ``None`` if botocore is unavailable or the service model cannot be
loaded, so the module still imports cleanly.
loaded.
"""
try:
from botocore.loaders import Loader
@ -977,15 +976,22 @@ def _load_bedrock_response_stream_shape():
return ServiceModel(service_dict).shape_for("ResponseStream")
except Exception as e:
verbose_logger.warning(
"litellm: could not pre-load bedrock-runtime response stream shape "
"litellm: could not load bedrock-runtime response stream shape "
"— Bedrock event-stream decoding will be unavailable. Error: %s",
e,
)
return None
# Eagerly resolved once per process — avoids per-instance or per-request disk I/O.
BEDROCK_RESPONSE_STREAM_SHAPE = _load_bedrock_response_stream_shape()
@functools.lru_cache(maxsize=1)
def get_bedrock_response_stream_shape():
"""
Lazily load and cache the bedrock-runtime ResponseStream shape for the process.
Avoids importing botocore (and logging warnings) unless Bedrock event-stream
decoding is actually needed.
"""
return _load_bedrock_response_stream_shape()
class BedrockEventStreamDecoderBase:
@ -999,7 +1005,8 @@ class BedrockEventStreamDecoderBase:
self.parser = EventStreamJSONParser()
def _parse_message_from_event(self, event) -> Optional[str]:
if BEDROCK_RESPONSE_STREAM_SHAPE is None:
response_stream_shape = get_bedrock_response_stream_shape()
if response_stream_shape is None:
raise BedrockError(
status_code=500,
message=(
@ -1008,9 +1015,7 @@ class BedrockEventStreamDecoderBase:
),
)
response_dict = event.to_response_dict()
parsed_response = self.parser.parse(
response_dict, BEDROCK_RESPONSE_STREAM_SHAPE
)
parsed_response = self.parser.parse(response_dict, response_stream_shape)
if response_dict["status_code"] != 200:
decoded_body = response_dict["body"].decode()

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

@ -45,6 +45,7 @@ from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import GenericStreamingChunk
from litellm.types.utils import GenericStreamingChunk as GChunk
from litellm.types.utils import ModelResponseStream
from litellm.utils import _supports_factory
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
@ -557,7 +558,29 @@ class AmazonAnthropicClaudeMessagesConfig(
anthropic_messages_request=anthropic_messages_request,
)
# 5a. Remove `custom` field from tools (Bedrock doesn't support it)
# 5a. Bedrock Invoke supports output_config (effort) for Claude 4.6+ models,
# but older models do not — strip it to avoid request rejection.
# Ref: https://github.com/BerriAI/litellm/issues/22797
if not (
_supports_factory(
model=model,
custom_llm_provider="bedrock",
key="supports_output_config",
)
or AnthropicConfig._model_supports_effort_param(model)
):
if anthropic_messages_request.pop("output_config", None) is not None:
verbose_logger.warning(
"Bedrock Invoke: stripping unsupported `output_config` for "
"model=%s — neither `supports_output_config` nor any "
"`supports_*_reasoning_effort` flag is set in "
"model_prices_and_context_window.json. Add the capability "
"flag to the model JSON entry if this model accepts "
"`output_config`.",
model,
)
# 5b. Remove `custom` field from tools (Bedrock doesn't support it)
# Claude Code sends `custom: {defer_loading: true}` on tool definitions,
# which causes Bedrock to reject the request with "Extra inputs are not permitted"
# Ref: https://github.com/BerriAI/litellm/issues/22847

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,7 +1,5 @@
import json
from typing import Any, Optional
from typing import Any, Dict, Optional
from litellm.constants import STREAM_SSE_DONE_STRING
from litellm.exceptions import AuthenticationError
from litellm.litellm_core_utils.core_helpers import process_response_headers
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
@ -9,13 +7,17 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
)
from litellm.llms.openai.common_utils import OpenAIError
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
from litellm.responses.sse_output_recovery import (
parse_sse_json_chunk,
record_output_item_chunk,
record_output_text_chunk,
)
from litellm.types.llms.openai import (
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
)
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders
from litellm.utils import CustomStreamWrapper
from ..authenticator import Authenticator
from ..common_utils import (
@ -111,86 +113,139 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig):
raw_response: Any,
logging_obj: Any,
):
content_type = (raw_response.headers or {}).get("content-type", "")
body_text = raw_response.text or ""
if "text/event-stream" not in content_type.lower():
trimmed_body = body_text.lstrip()
if not (
trimmed_body.startswith("event:")
or trimmed_body.startswith("data:")
or "\nevent:" in body_text
or "\ndata:" in body_text
):
return super().transform_response_api_response(
model=model,
raw_response=raw_response,
logging_obj=logging_obj,
)
if not self._should_parse_as_sse(
raw_response=raw_response, body_text=body_text
):
return super().transform_response_api_response(
model=model,
raw_response=raw_response,
logging_obj=logging_obj,
)
logging_obj.post_call(
original_response=raw_response.text,
additional_args={"complete_input_dict": {}},
)
completed_response = None
error_message = None
for chunk in body_text.splitlines():
stripped_chunk = CustomStreamWrapper._strip_sse_data_from_chunk(chunk)
if not stripped_chunk:
continue
stripped_chunk = stripped_chunk.strip()
if not stripped_chunk:
continue
if stripped_chunk == STREAM_SSE_DONE_STRING:
break
try:
parsed_chunk = json.loads(stripped_chunk)
except json.JSONDecodeError:
continue
if not isinstance(parsed_chunk, dict):
continue
event_type = parsed_chunk.get("type")
if event_type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED:
response_payload = parsed_chunk.get("response")
if isinstance(response_payload, dict):
response_payload = dict(response_payload)
if "created_at" in response_payload:
response_payload["created_at"] = _safe_convert_created_field(
response_payload["created_at"]
)
try:
completed_response = ResponsesAPIResponse(**response_payload)
except Exception:
completed_response = ResponsesAPIResponse.model_construct(
**response_payload
)
break
if event_type in (
ResponsesAPIStreamEvents.RESPONSE_FAILED,
ResponsesAPIStreamEvents.ERROR,
):
error_obj = parsed_chunk.get("error") or (
parsed_chunk.get("response") or {}
).get("error")
if error_obj is not None:
if isinstance(error_obj, dict):
error_message = error_obj.get("message") or str(error_obj)
else:
error_message = str(error_obj)
completed_response, error_message = self._extract_completed_response_from_sse(
body_text=body_text
)
if completed_response is None:
raise OpenAIError(
message=error_message or raw_response.text,
status_code=raw_response.status_code,
)
self._attach_response_headers(
completed_response=completed_response, raw_response=raw_response
)
return completed_response
def _should_parse_as_sse(self, raw_response: Any, body_text: str) -> bool:
content_type = (raw_response.headers or {}).get("content-type", "")
if "text/event-stream" in content_type.lower():
return True
trimmed_body = body_text.lstrip()
return bool(
trimmed_body.startswith("event:")
or trimmed_body.startswith("data:")
or "\nevent:" in body_text
or "\ndata:" in body_text
)
def _extract_completed_response_from_sse(
self, body_text: str
) -> tuple[Optional[ResponsesAPIResponse], Optional[str]]:
completed_response = None
error_message = None
streamed_output_items: Dict[int, dict] = {}
text_only_output_items: Dict[int, dict] = {}
for chunk in body_text.splitlines():
parsed_chunk = parse_sse_json_chunk(chunk)
if parsed_chunk is None:
continue
event_type = parsed_chunk.get("type")
if event_type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE:
record_output_item_chunk(
parsed_chunk=parsed_chunk,
output_items=streamed_output_items,
)
continue
if event_type == ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE:
record_output_text_chunk(
parsed_chunk=parsed_chunk,
output_items=streamed_output_items,
text_only_items=text_only_output_items,
)
continue
if event_type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED:
# Real OUTPUT_ITEM_DONE events take precedence at any given
# output_index, but text-only items at indices without a
# matching OUTPUT_ITEM_DONE must still be preserved (e.g.
# providers that emit only OUTPUT_TEXT_DONE for some indices).
merged_items: Dict[int, dict] = {**text_only_output_items}
merged_items.update(streamed_output_items)
completed_response = self._build_completed_response_from_chunk(
parsed_chunk=parsed_chunk,
streamed_output_items=merged_items,
)
break
if event_type in (
ResponsesAPIStreamEvents.RESPONSE_FAILED,
ResponsesAPIStreamEvents.ERROR,
):
extracted_error = self._extract_error_message(parsed_chunk)
if extracted_error is not None:
error_message = extracted_error
return completed_response, error_message
def _build_completed_response_from_chunk(
self, parsed_chunk: Dict[str, Any], streamed_output_items: Dict[int, dict]
) -> Optional[ResponsesAPIResponse]:
response_payload = parsed_chunk.get("response")
if not isinstance(response_payload, dict):
return None
response_payload = dict(response_payload)
if not response_payload.get("output") and streamed_output_items:
response_payload["output"] = [
item for _, item in sorted(streamed_output_items.items())
]
if "created_at" in response_payload:
response_payload["created_at"] = _safe_convert_created_field(
response_payload["created_at"]
)
try:
return ResponsesAPIResponse(**response_payload)
except Exception:
return ResponsesAPIResponse.model_construct(**response_payload)
def _extract_error_message(self, parsed_chunk: Dict[str, Any]) -> Optional[str]:
error_obj = parsed_chunk.get("error") or (
parsed_chunk.get("response") or {}
).get("error")
if error_obj is None:
return None
if isinstance(error_obj, dict):
return error_obj.get("message") or str(error_obj)
return str(error_obj)
def _attach_response_headers(
self,
completed_response: ResponsesAPIResponse,
raw_response: Any,
) -> None:
raw_headers = dict(raw_response.headers)
processed_headers = process_response_headers(raw_headers)
if not hasattr(completed_response, "_hidden_params"):
setattr(completed_response, "_hidden_params", {})
completed_response._hidden_params["additional_headers"] = processed_headers
completed_response._hidden_params["headers"] = raw_headers
return completed_response
def get_complete_url(
self,

View file

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

View file

@ -257,14 +257,19 @@ class GenericContainerHandler:
returns_binary = endpoint_config.get("returns_binary", False)
is_multipart = endpoint_config.get("is_multipart", False)
# An empty dict passed as `params` to httpx strips any existing query
# string from the URL (e.g. ?api-version=...). Use None instead so
# httpx leaves the URL's own query string intact.
effective_params = query_params or None
try:
if method == "GET":
response = http_client.get(
url=url, headers=headers, params=query_params
url=url, headers=headers, params=effective_params
)
elif method == "DELETE":
response = http_client.delete(
url=url, headers=headers, params=query_params
url=url, headers=headers, params=effective_params
)
elif method == "POST":
if is_multipart and "file" in kwargs:
@ -272,11 +277,11 @@ class GenericContainerHandler:
kwargs["file"], headers
)
response = http_client.post(
url=url, headers=headers, params=query_params, files=files
url=url, headers=headers, params=effective_params, files=files
)
else:
response = http_client.post(
url=url, headers=headers, params=query_params
url=url, headers=headers, params=effective_params
)
else:
raise ValueError(f"Unsupported HTTP method: {method}")
@ -376,14 +381,19 @@ class GenericContainerHandler:
returns_binary = endpoint_config.get("returns_binary", False)
is_multipart = endpoint_config.get("is_multipart", False)
# An empty dict passed as `params` to httpx strips any existing query
# string from the URL (e.g. ?api-version=...). Use None instead so
# httpx leaves the URL's own query string intact.
effective_params = query_params or None
try:
if method == "GET":
response = await http_client.get(
url=url, headers=headers, params=query_params
url=url, headers=headers, params=effective_params
)
elif method == "DELETE":
response = await http_client.delete(
url=url, headers=headers, params=query_params
url=url, headers=headers, params=effective_params
)
elif method == "POST":
if is_multipart and "file" in kwargs:
@ -391,11 +401,11 @@ class GenericContainerHandler:
kwargs["file"], headers
)
response = await http_client.post(
url=url, headers=headers, params=query_params, files=files
url=url, headers=headers, params=effective_params, files=files
)
else:
response = await http_client.post(
url=url, headers=headers, params=query_params
url=url, headers=headers, params=effective_params
)
else:
raise ValueError(f"Unsupported HTTP method: {method}")

View file

@ -1409,6 +1409,8 @@ class BaseLLMHTTPHandler:
document=document,
optional_params=optional_params,
headers=headers,
api_key=api_key,
api_base=api_base,
)
# All providers return OCRRequestData
@ -1477,6 +1479,8 @@ class BaseLLMHTTPHandler:
document=document,
optional_params=optional_params,
headers=headers,
api_key=api_key,
api_base=api_base,
)
# All providers return OCRRequestData
@ -7834,7 +7838,7 @@ class BaseLLMHTTPHandler:
response = sync_httpx_client.get(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_list_response(
@ -7911,7 +7915,7 @@ class BaseLLMHTTPHandler:
response = await async_httpx_client.get(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_list_response(
@ -8001,7 +8005,7 @@ class BaseLLMHTTPHandler:
response = sync_httpx_client.get(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_retrieve_response(
@ -8078,7 +8082,7 @@ class BaseLLMHTTPHandler:
response = await async_httpx_client.get(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_retrieve_response(
@ -8168,7 +8172,7 @@ class BaseLLMHTTPHandler:
response = sync_httpx_client.delete(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_delete_response(
@ -8245,7 +8249,7 @@ class BaseLLMHTTPHandler:
response = await async_httpx_client.delete(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_delete_response(
@ -8341,7 +8345,7 @@ class BaseLLMHTTPHandler:
response = sync_httpx_client.get(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_file_list_response(
@ -8420,7 +8424,7 @@ class BaseLLMHTTPHandler:
response = await async_httpx_client.get(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_file_list_response(
@ -8508,7 +8512,7 @@ class BaseLLMHTTPHandler:
response = sync_httpx_client.get(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_file_content_response(
@ -8584,7 +8588,7 @@ class BaseLLMHTTPHandler:
response = await async_httpx_client.get(
url=url,
headers=headers,
params=params,
params=params or None,
)
return container_provider_config.transform_container_file_content_response(

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

@ -4,6 +4,7 @@ from typing import Any, List, Literal, Optional, Tuple, Union, cast
import httpx
import litellm
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -26,6 +27,7 @@ from litellm.types.utils import (
ProviderSpecificModelInfo,
)
from litellm.utils import (
get_model_cost_mutation_generation,
supports_function_calling,
supports_reasoning,
supports_tool_choice,
@ -112,6 +114,19 @@ class FireworksAIConfig(OpenAIGPTConfig):
# Only add tools for models that support function calling
if supports_function_calling(model=model, custom_llm_provider="fireworks_ai"):
supported_params.append("tools")
supported_params.append("parallel_tool_calls")
else:
# Historically every Fireworks model advertised tool support, so a
# JSON entry that flips `supports_function_calling` to false will
# silently drop `tools` from requests. Surface this so users can
# tell why their tool calls suddenly stop working.
verbose_logger.debug(
"fireworks_ai model %r is marked as not supporting "
"function calling in model_prices_and_context_window.json; "
"`tools` and `parallel_tool_calls` will be dropped from the "
"request.",
model,
)
# Only add tool_choice for models that explicitly support it
if supports_tool_choice(model=model, custom_llm_provider="fireworks_ai"):
@ -251,34 +266,100 @@ class FireworksAIConfig(OpenAIGPTConfig):
return messages
def get_provider_info(self, model: str) -> ProviderSpecificModelInfo:
# Models that support reasoning_effort
reasoning_supported_models = [
"qwen3-8b",
"qwen3-32b",
"qwen3-coder-480b-a35b-instruct",
"deepseek-v3p1",
"deepseek-v3p2",
"glm-4p5",
"glm-4p5-air",
"glm-4p6",
"gpt-oss-120b",
"gpt-oss-20b",
# Cached index of fireworks_ai/* entries from litellm.model_cost. Building
# this index requires a full scan of model_cost (tens of thousands of
# entries), so we memoize it. The cache key is (id(model_cost),
# mutation_generation): the generation counter is bumped on every
# register_model / reload path, so add+remove or in-place value
# replacement (which can leave id and len unchanged) still invalidates.
_fireworks_index_cache: Optional[Tuple[int, int, List[Tuple[str, dict]]]] = None
@classmethod
def _get_fireworks_index(cls) -> List[Tuple[str, dict]]:
model_cost = litellm.model_cost
signature = (id(model_cost), get_model_cost_mutation_generation())
cached = cls._fireworks_index_cache
if (
cached is not None
and cached[0] == signature[0]
and cached[1] == signature[1]
):
return cached[2]
index: List[Tuple[str, dict]] = []
for key, model_info in model_cost.items():
if not key.startswith("fireworks_ai/"):
continue
if not isinstance(model_info, dict):
continue
key_short = key[len("fireworks_ai/") :]
if key_short.startswith("accounts/fireworks/models/"):
key_short = key_short[len("accounts/fireworks/models/") :]
if not key_short:
continue
index.append((key_short, model_info))
cls._fireworks_index_cache = (signature[0], signature[1], index)
return index
@staticmethod
def _matches_on_hyphen_boundary(short_name: str, key_short: str) -> bool:
"""Return True if `key_short` appears in `short_name` aligned to
hyphen-separated word boundaries (or end-of-string). This avoids
spurious substring matches like `"some-model"` matching
`"awesome-model"`."""
if short_name == key_short:
return True
if short_name.startswith(key_short + "-"):
return True
if short_name.endswith("-" + key_short):
return True
return ("-" + key_short + "-") in short_name
def _get_model_cost_capability(self, model: str, capability: str) -> Optional[bool]:
short_name = model
if short_name.startswith("fireworks_ai/"):
short_name = short_name[len("fireworks_ai/") :]
if short_name.startswith("accounts/fireworks/models/"):
short_name = short_name[len("accounts/fireworks/models/") :]
candidate_keys = [
model,
f"fireworks_ai/{short_name}",
f"fireworks_ai/accounts/fireworks/models/{short_name}",
]
# Normalize model name - remove prefix if present
normalized_model = model
if model.startswith("fireworks_ai/"):
normalized_model = model.replace("fireworks_ai/", "")
if normalized_model.startswith("accounts/fireworks/models/"):
normalized_model = normalized_model.replace(
"accounts/fireworks/models/", ""
)
for candidate_key in candidate_keys:
model_info = litellm.model_cost.get(candidate_key)
if model_info is not None and model_info.get(capability) is not None:
return cast(Optional[bool], model_info.get(capability))
# Check if model supports reasoning
supports_reasoning_value = any(
reasoning_model in normalized_model
for reasoning_model in reasoning_supported_models
# Fallback: preserve historical substring matching for model name
# variants (e.g. fine-tuned or regionally-suffixed versions of a
# known model). Pick the *longest* matching entry so a more specific
# known model (e.g. "qwen3-8b-instruct") wins over a less specific
# one (e.g. "qwen3-8b") when the query model is more specific still.
# Use hyphen-aligned matching to avoid false positives where a short
# known model name is an unrelated substring of a longer one.
best_match_short: Optional[str] = None
best_match_value: Optional[bool] = None
for key_short, model_info in self._get_fireworks_index():
if model_info.get(capability) is None:
continue
if not self._matches_on_hyphen_boundary(short_name, key_short):
continue
if best_match_short is None or len(key_short) > len(best_match_short):
best_match_short = key_short
best_match_value = cast(Optional[bool], model_info.get(capability))
return best_match_value
def get_provider_info(self, model: str) -> ProviderSpecificModelInfo:
supports_function_calling_value = self._get_model_cost_capability(
model=model, capability="supports_function_calling"
)
supports_reasoning_value = self._get_model_cost_capability(
model=model, capability="supports_reasoning"
)
provider_specific_model_info: ProviderSpecificModelInfo = {
@ -288,9 +369,16 @@ class FireworksAIConfig(OpenAIGPTConfig):
"supports_vision": True, # via document inlining
}
if supports_function_calling_value is not None:
provider_specific_model_info["supports_function_calling"] = (
supports_function_calling_value
)
# Only include supports_reasoning if True
if supports_reasoning_value:
provider_specific_model_info["supports_reasoning"] = True
provider_specific_model_info["supports_reasoning"] = (
supports_reasoning_value
)
return provider_specific_model_info

View file

View file

@ -0,0 +1,298 @@
"""
Google AI Studio Agents API configuration.
Proxies the Gemini v1beta Agents API:
POST /v1beta/agents create
GET /v1beta/agents list
GET /v1beta/agents/{name} get
DELETE /v1beta/agents/{name} delete
GET /v1beta/agents/{name}/versions list versions
"""
from typing import Any, Dict, Optional, Tuple, Union
import httpx
from litellm._logging import verbose_logger
from litellm.llms.base_llm.agents.transformation import BaseAgentsAPIConfig
from litellm.llms.gemini.common_utils import GeminiError, GeminiModelInfo
from litellm.types.agents import (
AgentCreateResponse,
AgentDeleteResult,
AgentListResponse,
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")
# LiteLLM-internal keys that must never be forwarded to Gemini.
_LITELLM_INTERNAL_KEYS = frozenset(
{
"custom_llm_provider",
"api_key",
"api_base",
"make_public",
"cost_per_query",
"input_cost_per_token",
"output_cost_per_token",
"require_trace_id_on_calls_to_agent",
"require_trace_id_on_calls_by_agent",
"max_iterations",
"max_budget_per_session",
"guardrails",
"is_public",
"agent_name",
"agent_id",
"agent_card_params",
"provider_agent_response",
}
)
class GeminiAgentsConfig(BaseAgentsAPIConfig):
"""
Configuration for the Google AI Studio (Gemini) native Agents API.
Authentication uses x-goog-api-key, resolved from (in order):
1. litellm_params["api_key"]
2. GOOGLE_API_KEY env var
3. GEMINI_API_KEY env var
"""
@property
def api_version(self) -> str:
return "v1beta"
def _base_url(self, api_base: Optional[str]) -> str:
return f"{GeminiModelInfo.get_api_base(api_base)}/{self.api_version}"
# ------------------------------------------------------------------ #
# Shared helpers #
# ------------------------------------------------------------------ #
def get_error_class(
self,
error_message: str,
status_code: int,
headers: Union[dict, httpx.Headers],
) -> Exception:
return GeminiError(
message=error_message,
status_code=status_code,
headers=dict(headers),
)
def get_complete_url(
self,
api_base: Optional[str],
litellm_params: Dict[str, Any],
) -> str:
return f"{self._base_url(api_base)}/agents"
def validate_environment(
self,
headers: Dict[str, str],
litellm_params: Dict[str, Any],
) -> Dict[str, str]:
headers = dict(headers)
headers["Content-Type"] = "application/json"
explicit_api_key = litellm_params.get("api_key")
# SECURITY: when the caller overrides ``api_base``, refuse to fall back
# to the process-wide GOOGLE_API_KEY / GEMINI_API_KEY env vars. Otherwise
# an authenticated proxy user could set ``api_base`` to an attacker-
# controlled host and have the proxy ship its shared Gemini key in the
# ``x-goog-api-key`` header.
if litellm_params.get("api_base") and not explicit_api_key:
raise ValueError(
"When overriding api_base for Gemini agents, you must also "
"supply an explicit api_key. Falling back to GOOGLE_API_KEY / "
"GEMINI_API_KEY env vars with a custom api_base is refused "
"to prevent leaking the shared provider key to arbitrary hosts."
)
api_key = GeminiModelInfo.get_api_key(explicit_api_key)
if not api_key:
raise ValueError(
"Google API key is required. "
"Set GOOGLE_API_KEY or GEMINI_API_KEY, or pass api_key."
)
headers["x-goog-api-key"] = api_key
return headers
def _raise_for_status(self, raw_response: httpx.Response) -> None:
if not (200 <= raw_response.status_code < 300):
raise GeminiError(
message=raw_response.text,
status_code=raw_response.status_code,
headers=dict(raw_response.headers),
)
# ------------------------------------------------------------------ #
# CREATE #
# ------------------------------------------------------------------ #
def transform_create_request(
self,
name: str,
litellm_params: Dict[str, Any],
) -> Dict[str, Any]:
body: Dict[str, Any] = {"name": name}
for key in _GEMINI_AGENT_BODY_KEYS:
value = litellm_params.get(key)
if value is not None:
body[key] = value
verbose_logger.debug("GeminiAgentsConfig create body: %s", body)
return body
def transform_create_response(
self,
raw_response: httpx.Response,
name: str,
) -> AgentCreateResponse:
"""
Gemini returns:
{"id": "my-agent", "base_agent": "waverunner",
"system_instruction": "...", "base_environment": {...}}
"""
self._raise_for_status(raw_response)
try:
data: Dict[str, Any] = raw_response.json()
except Exception:
verbose_logger.warning(
"GeminiAgentsConfig: non-JSON create response (status=%d).",
raw_response.status_code,
)
data = {"id": name}
# Gemini uses "id" as the identifier; normalise to both fields.
data.setdefault("id", name)
data.setdefault("name", data["id"])
verbose_logger.debug("GeminiAgentsConfig create response: %s", data)
return AgentCreateResponse(**data)
# ------------------------------------------------------------------ #
# LIST #
# ------------------------------------------------------------------ #
def transform_list_request(
self,
api_base: Optional[str],
litellm_params: Dict[str, Any],
) -> Tuple[str, Dict[str, Any]]:
url = f"{self._base_url(api_base)}/agents"
params: Dict[str, Any] = {}
if litellm_params.get("page_size"):
params["pageSize"] = litellm_params["page_size"]
if litellm_params.get("page_token"):
params["pageToken"] = litellm_params["page_token"]
return url, params
def transform_list_response(
self,
raw_response: httpx.Response,
) -> AgentListResponse:
self._raise_for_status(raw_response)
try:
data = raw_response.json()
except Exception:
data = {}
verbose_logger.debug("GeminiAgentsConfig list response: %s", data)
return AgentListResponse(
agents=data.get("agents", []),
next_page_token=data.get("nextPageToken"),
)
# ------------------------------------------------------------------ #
# GET #
# ------------------------------------------------------------------ #
def transform_get_request(
self,
name: str,
api_base: Optional[str],
litellm_params: Dict[str, Any],
) -> Tuple[str, Dict[str, Any]]:
url = f"{self._base_url(api_base)}/agents/{name}"
return url, {}
def transform_get_response(
self,
raw_response: httpx.Response,
name: str,
) -> AgentCreateResponse:
"""Same shape as create response — Gemini returns "id" as identifier."""
self._raise_for_status(raw_response)
try:
data = raw_response.json()
except Exception:
data = {"id": name}
data.setdefault("id", name)
data.setdefault("name", data["id"])
verbose_logger.debug("GeminiAgentsConfig get response: %s", data)
return AgentCreateResponse(**data)
# ------------------------------------------------------------------ #
# DELETE #
# ------------------------------------------------------------------ #
def transform_delete_request(
self,
name: str,
api_base: Optional[str],
litellm_params: Dict[str, Any],
) -> str:
return f"{self._base_url(api_base)}/agents/{name}"
def transform_delete_response(
self,
raw_response: httpx.Response,
name: str,
) -> AgentDeleteResult:
"""Gemini returns an empty body ``{}`` with HTTP 200 on success."""
self._raise_for_status(raw_response)
verbose_logger.debug(
"GeminiAgentsConfig delete (status=%d) agent '%s'",
raw_response.status_code,
name,
)
return AgentDeleteResult(name=name, deleted=True)
# ------------------------------------------------------------------ #
# LIST VERSIONS #
# ------------------------------------------------------------------ #
def transform_list_versions_request(
self,
name: str,
api_base: Optional[str],
litellm_params: Dict[str, Any],
) -> Tuple[str, Dict[str, Any]]:
url = f"{self._base_url(api_base)}/agents/{name}/versions"
params: Dict[str, Any] = {}
if litellm_params.get("page_size"):
params["pageSize"] = litellm_params["page_size"]
if litellm_params.get("page_token"):
params["pageToken"] = litellm_params["page_token"]
return url, params
def transform_list_versions_response(
self,
raw_response: httpx.Response,
name: str,
) -> AgentVersionsResponse:
"""
Gemini returns:
{"agentVersions": [{"agent": "waverunner", "name": "agents/.../versions/uuid", ...}]}
"""
self._raise_for_status(raw_response)
try:
data = raw_response.json()
except Exception:
data = {}
verbose_logger.debug(
"GeminiAgentsConfig list_versions response for '%s': %s", name, data
)
return AgentVersionsResponse(
agent_versions=data.get("agentVersions", []),
next_page_token=data.get("nextPageToken"),
)

View file

@ -164,5 +164,8 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig):
# If conversion fails, leave as is and let the API handle it
pass
return _gemini_convert_messages_with_history(
messages=messages, model=model, litellm_params=litellm_params
messages=messages,
model=model,
litellm_params=litellm_params,
custom_llm_provider="gemini",
)

View file

@ -6,13 +6,18 @@ Per OpenAPI spec (https://ai.google.dev/static/api/interactions.openapi.json):
- Get: GET https://generativelanguage.googleapis.com/{api_version}/interactions/{interaction_id}
- Delete: DELETE https://generativelanguage.googleapis.com/{api_version}/interactions/{interaction_id}
This is a thin wrapper - no transformation needed since we follow the spec directly.
Schema versioning:
- Default (Api-Revision: 2026-05-20): new `steps` schema.
- Legacy (Api-Revision: 2026-05-07): old `outputs` schema, controlled via
litellm.use_legacy_interactions_schema = True. Remove flag after June 8, 2026.
"""
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.core_helpers import process_response_headers
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
@ -64,6 +69,7 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig):
"stream",
"store",
"background",
"environment",
"response_modalities",
"response_format",
"response_mime_type",
@ -83,6 +89,15 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig):
api_key = GeminiModelInfo.get_api_key(litellm_params.get("api_key"))
if api_key:
headers["x-goog-api-key"] = api_key
# Inject the Api-Revision header to select the response schema.
# Default to the new `steps` schema unless the operator has opted out.
# Remove this conditional after June 8, 2026 and always use 2026-05-20.
if litellm.use_legacy_interactions_schema:
headers["Api-Revision"] = "2026-05-07"
else:
headers["Api-Revision"] = "2026-05-20"
return headers
def get_complete_url(
@ -118,8 +133,19 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig):
headers: dict,
) -> Dict:
"""
Build request body per OpenAPI spec - minimal transformation.
Build request body per OpenAPI spec.
When on the new schema (use_legacy_interactions_schema=False, the default):
- ``response_mime_type`` is folded into ``response_format`` and stripped from
the body (the field was removed in Api-Revision 2026-05-20).
- ``generation_config.image_config`` is moved to a ``response_format`` entry
with ``"type": "image"`` (also removed from generation_config in 2026-05-20).
When on the legacy schema (use_legacy_interactions_schema=True):
- All fields are forwarded as-is.
"""
use_legacy: bool = litellm.use_legacy_interactions_schema
request_body: Dict[str, Any] = {}
# Model or Agent (one required)
@ -134,23 +160,81 @@ class GoogleAIStudioInteractionsConfig(BaseInteractionsAPIConfig):
if input is not None:
request_body["input"] = input
# Pass through optional params directly (they match the spec)
# Pass through optional params — legacy schema keeps all fields as-is.
optional_keys = [
"tools",
"system_instruction",
"generation_config",
"stream",
"store",
"background",
"environment",
"response_modalities",
"response_format",
"response_mime_type",
"previous_interaction_id",
]
for key in optional_keys:
if optional_params.get(key) is not None:
request_body[key] = optional_params[key]
if use_legacy:
# Legacy schema: forward response_mime_type and response_format as-is.
for key in ("response_format", "response_mime_type", "generation_config"):
if optional_params.get(key) is not None:
request_body[key] = optional_params[key]
else:
# New schema (Api-Revision: 2026-05-20):
# response_mime_type is removed — fold it into response_format.
response_format = optional_params.get("response_format")
response_mime_type = optional_params.get("response_mime_type")
if (
response_mime_type
and not isinstance(response_format, list)
and (
not isinstance(response_format, dict)
or "mime_type" not in response_format
)
):
# Wrap the legacy schema into the new polymorphic format.
new_rf: Dict[str, Any] = {
"type": "text",
"mime_type": response_mime_type,
}
if response_format is not None:
new_rf["schema"] = response_format
response_format = new_rf
if response_format is not None:
request_body["response_format"] = response_format
# image_config moves out of generation_config into response_format.
generation_config: Optional[Dict[str, Any]] = optional_params.get(
"generation_config"
)
if generation_config is not None:
image_config = None
if isinstance(generation_config, dict):
generation_config = dict(
generation_config
) # avoid mutating the caller's dict
image_config = generation_config.pop("image_config", None)
if not generation_config:
generation_config = None
if generation_config is not None:
request_body["generation_config"] = generation_config
if image_config is not None:
# Move image_config to response_format with type=image.
image_rf: Dict[str, Any] = {"type": "image", **image_config}
existing_rf = request_body.get("response_format")
if existing_rf is None:
request_body["response_format"] = image_rf
elif isinstance(existing_rf, list):
request_body["response_format"] = [*existing_rf, image_rf]
else:
# Convert single entry to array for multimodal output.
request_body["response_format"] = [existing_rf, image_rf]
return request_body
def transform_response(

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

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

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

@ -0,0 +1 @@

View file

@ -0,0 +1,159 @@
import base64
import binascii
from collections import defaultdict
from typing import TYPE_CHECKING, Any, Dict, List, NoReturn, Optional, Tuple
from litellm.constants import request_timeout
REDUCTO_API_BASE = "https://platform.reducto.ai"
REDUCTO_ID_PREFIX = "reducto://"
if TYPE_CHECKING:
from litellm.llms.base_llm.ocr.transformation import OCRPage
def _normalize_api_base(api_base: Optional[str]) -> str:
return (api_base or REDUCTO_API_BASE).rstrip("/")
def _raise_bad_request(message: str, model: str) -> NoReturn:
import litellm
raise litellm.BadRequestError(
message=message,
model=model,
llm_provider="reducto",
)
def extract_file_id_or_bytes(
source_url: str,
model: str,
) -> Tuple[Optional[str], Optional[bytes], Optional[str]]:
if source_url.startswith(REDUCTO_ID_PREFIX):
return source_url, None, None
if source_url.startswith("http://") or source_url.startswith("https://"):
_raise_bad_request(
"Reducto requires type='file' (auto-uploaded) or a reducto:// id. Plain http(s) URLs are not supported; upload the file first.",
model=model,
)
if not source_url.startswith("data:"):
_raise_bad_request(
"Reducto requires a reducto:// id or a base64 data URI after OCR preprocessing.",
model=model,
)
try:
header, encoded = source_url.split(",", 1)
except ValueError:
_raise_bad_request("Invalid Reducto data URI provided.", model=model)
if ";base64" not in header:
_raise_bad_request(
"Reducto only supports base64-encoded data URIs.", model=model
)
mime = header.removeprefix("data:").split(";")[0] or "application/octet-stream"
try:
raw_bytes = base64.b64decode(encoded, validate=True)
except (binascii.Error, ValueError):
_raise_bad_request("Invalid Reducto base64 payload provided.", model=model)
return None, raw_bytes, mime
def _extract_file_id_from_upload_response(response: Any) -> str:
try:
payload = response.json()
except ValueError as exc:
raise ValueError(
"Reducto /upload returned a non-JSON 200 response: {}".format(response.text)
) from exc
file_id = (payload or {}).get("file_id") if isinstance(payload, dict) else None
if not isinstance(file_id, str) or not file_id:
raise ValueError(
"Reducto /upload returned 200 without a file_id; got payload={}".format(
payload
)
)
return file_id
def upload_bytes_sync(
raw_bytes: bytes,
mime: Optional[str],
api_key: str,
api_base: Optional[str],
) -> str:
import litellm
response = litellm.module_level_client.post(
url="{}{}".format(_normalize_api_base(api_base), "/upload"),
headers={"Authorization": f"Bearer {api_key}"},
files={"file": ("document", raw_bytes, mime or "application/octet-stream")},
timeout=request_timeout,
)
response.raise_for_status()
return _extract_file_id_from_upload_response(response)
async def upload_bytes_async(
raw_bytes: bytes,
mime: Optional[str],
api_key: str,
api_base: Optional[str],
) -> str:
import litellm
response = await litellm.module_level_aclient.post(
url="{}{}".format(_normalize_api_base(api_base), "/upload"),
headers={"Authorization": f"Bearer {api_key}"},
files={"file": ("document", raw_bytes, mime or "application/octet-stream")},
timeout=request_timeout,
)
response.raise_for_status()
return _extract_file_id_from_upload_response(response)
def build_pages_from_reducto(result: Dict[str, Any]) -> List["OCRPage"]:
from litellm.llms.base_llm.ocr.transformation import OCRPage
chunks = result.get("chunks", []) or []
blocks_by_page: Dict[int, List[Dict[str, Any]]] = defaultdict(list)
for chunk in chunks:
for block in chunk.get("blocks", []) or []:
page_no = (block.get("bbox") or {}).get("page")
if page_no is None:
continue
try:
normalized_page = int(page_no)
except (TypeError, ValueError):
continue
blocks_by_page[normalized_page].append(block)
if not blocks_by_page:
fallback_markdown = "\n\n".join(
chunk.get("content", "") for chunk in chunks if chunk.get("content")
)
if fallback_markdown == "":
return []
return [OCRPage(index=0, markdown=fallback_markdown)]
pages: List["OCRPage"] = []
for page_no, blocks in sorted(blocks_by_page.items()):
markdown = "\n\n".join(
block.get("content", "") for block in blocks if block.get("content")
)
page_index = max(page_no - 1, 0)
page = OCRPage(
index=page_index,
markdown=markdown,
)
# OCRPage accepts extra keys at runtime; assign blocks after construction
# so static typing does not reject provider-specific metadata.
setattr(page, "blocks", blocks)
pages.append(page)
return pages

View file

@ -0,0 +1 @@

View file

@ -0,0 +1,241 @@
from typing import Any, Dict, Optional, Tuple
import httpx
from litellm.llms.base_llm.ocr.transformation import (
BaseOCRConfig,
DocumentType,
OCRRequestData,
OCRResponse,
OCRUsageInfo,
)
from litellm.llms.reducto.common import (
REDUCTO_API_BASE,
build_pages_from_reducto,
extract_file_id_or_bytes,
upload_bytes_async,
upload_bytes_sync,
)
class _BaseReductoOCRConfig(BaseOCRConfig):
def map_ocr_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
) -> dict:
mapped_params = dict(optional_params)
supported_params = self.get_supported_ocr_params(model=model)
for param, value in non_default_params.items():
if param in supported_params:
mapped_params[param] = value
return mapped_params
def validate_environment(
self,
headers: Dict,
model: str,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
litellm_params: Optional[dict] = None,
**kwargs,
) -> Dict:
from litellm.secret_managers.main import get_secret_str
resolved_key = api_key or get_secret_str("REDUCTO_API_KEY")
if resolved_key is None:
raise ValueError(
"Missing REDUCTO_API_KEY - set it in the environment or pass api_key to litellm.ocr()/litellm.aocr()"
)
return {
"Authorization": f"Bearer {resolved_key}",
"Content-Type": "application/json",
**headers,
}
def get_complete_url(
self,
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: Optional[dict] = None,
**kwargs,
) -> str:
return "{}/parse".format((api_base or REDUCTO_API_BASE).rstrip("/"))
def _get_source_url(self, document: DocumentType, model: str) -> str:
source_url = document.get("document_url") or document.get("image_url")
if source_url is None:
raise ValueError(
"Reducto expected OCR preprocessing to produce document_url or image_url for model={}".format(
model
)
)
return source_url
@staticmethod
def _resolve_credentials(
api_key: Optional[str], api_base: Optional[str]
) -> Tuple[str, str]:
from litellm.secret_managers.main import get_secret_str
resolved_key = api_key or get_secret_str("REDUCTO_API_KEY")
if resolved_key is None:
raise ValueError(
"Missing REDUCTO_API_KEY - set it in the environment or pass api_key to litellm.ocr()/litellm.aocr()"
)
resolved_base = (api_base or REDUCTO_API_BASE).rstrip("/")
return resolved_key, resolved_base
def _ensure_file_id_sync(
self,
model: str,
document: DocumentType,
api_key: Optional[str],
api_base: Optional[str],
) -> str:
source_url = self._get_source_url(document=document, model=model)
file_id, raw_bytes, mime = extract_file_id_or_bytes(source_url, model=model)
if file_id is not None:
return file_id
resolved_key, resolved_base = self._resolve_credentials(api_key, api_base)
return upload_bytes_sync(
raw_bytes=raw_bytes or b"",
mime=mime,
api_key=resolved_key,
api_base=resolved_base,
)
async def _ensure_file_id_async(
self,
model: str,
document: DocumentType,
api_key: Optional[str],
api_base: Optional[str],
) -> str:
source_url = self._get_source_url(document=document, model=model)
file_id, raw_bytes, mime = extract_file_id_or_bytes(source_url, model=model)
if file_id is not None:
return file_id
resolved_key, resolved_base = self._resolve_credentials(api_key, api_base)
return await upload_bytes_async(
raw_bytes=raw_bytes or b"",
mime=mime,
api_key=resolved_key,
api_base=resolved_base,
)
def transform_ocr_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: Any,
**kwargs,
) -> OCRResponse:
response_json = raw_response.json()
result = response_json.get("result", response_json) or {}
usage = response_json.get("usage", {}) or {}
response = OCRResponse(
pages=build_pages_from_reducto(result),
model=model,
usage_info=OCRUsageInfo(
pages_processed=usage.get("num_pages"),
credits=usage.get("credits"),
),
object="ocr",
)
response._hidden_params["reducto_raw"] = response_json
return response
class ReductoParseV3Config(_BaseReductoOCRConfig):
def get_supported_ocr_params(self, model: str) -> list:
return ["formatting", "retrieval", "settings"]
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
file_id = self._ensure_file_id_sync(
model=model,
document=document,
api_key=kwargs.get("api_key"),
api_base=kwargs.get("api_base"),
)
return OCRRequestData(data={"input": file_id, **optional_params}, files=None)
async def async_transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
file_id = await self._ensure_file_id_async(
model=model,
document=document,
api_key=kwargs.get("api_key"),
api_base=kwargs.get("api_base"),
)
return OCRRequestData(data={"input": file_id, **optional_params}, files=None)
class ReductoParseLegacyConfig(_BaseReductoOCRConfig):
def get_supported_ocr_params(self, model: str) -> list:
return ["enhance"]
def _build_legacy_body(self, file_id: str, optional_params: dict) -> Dict[str, Any]:
body: Dict[str, Any] = {"document_url": file_id}
enhance = optional_params.get("enhance")
if enhance is not None:
body["options"] = {"enhance": enhance}
return body
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
file_id = self._ensure_file_id_sync(
model=model,
document=document,
api_key=kwargs.get("api_key"),
api_base=kwargs.get("api_base"),
)
return OCRRequestData(
data=self._build_legacy_body(
file_id=file_id, optional_params=optional_params
),
files=None,
)
async def async_transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
file_id = await self._ensure_file_id_async(
model=model,
document=document,
api_key=kwargs.get("api_key"),
api_base=kwargs.get("api_base"),
)
return OCRRequestData(
data=self._build_legacy_body(
file_id=file_id, optional_params=optional_params
),
files=None,
)

View file

@ -1,3 +1,4 @@
import functools
import json
from typing import AsyncIterator, Iterator, List, Optional, Union
@ -22,14 +23,22 @@ def _load_sagemaker_response_stream_shape():
)
except Exception as e:
verbose_logger.warning(
"litellm: could not pre-load sagemaker-runtime response stream shape "
"litellm: could not load sagemaker-runtime response stream shape "
"— SageMaker event-stream decoding will be unavailable. Error: %s",
e,
)
return None
SAGEMAKER_RESPONSE_STREAM_SHAPE = _load_sagemaker_response_stream_shape()
@functools.lru_cache(maxsize=1)
def get_sagemaker_response_stream_shape():
"""
Lazily load and cache the sagemaker-runtime stream shape for the process.
Avoids importing botocore (and logging warnings) unless SageMaker event-stream
decoding is actually needed.
"""
return _load_sagemaker_response_stream_shape()
class SagemakerError(BaseLLMException):
@ -207,7 +216,8 @@ class AWSEventStreamDecoder:
verbose_logger.error(f"Final error parsing accumulated JSON: {e}")
def _parse_message_from_event(self, event) -> Optional[str]:
if SAGEMAKER_RESPONSE_STREAM_SHAPE is None:
response_stream_shape = get_sagemaker_response_stream_shape()
if response_stream_shape is None:
raise SagemakerError(
status_code=500,
message=(
@ -216,9 +226,7 @@ class AWSEventStreamDecoder:
),
)
response_dict = event.to_response_dict()
parsed_response = self.parser.parse(
response_dict, SAGEMAKER_RESPONSE_STREAM_SHAPE
)
parsed_response = self.parser.parse(response_dict, response_stream_shape)
if response_dict["status_code"] != 200:
raise ValueError(f"Bad response code, expected 200: {response_dict}")

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

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

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.
@ -174,7 +174,9 @@ def transform_openai_messages_to_gemini_context_caching(
)
transformed_messages = _gemini_convert_messages_with_history(
messages=new_messages, model=model
messages=new_messages,
model=model,
custom_llm_provider=custom_llm_provider,
)
model_name = "models/{}".format(model)

View file

@ -41,7 +41,7 @@ class ContextCachingEndpoints(VertexBase):
"""
def __init__(self) -> None:
pass
super().__init__()
def _get_token_and_url_context_caching(
self,

View file

@ -682,6 +682,7 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915
messages: List[AllMessageValues],
model: Optional[str] = None,
litellm_params: Optional[dict] = None,
custom_llm_provider: Optional[str] = None,
) -> List[ContentType]:
"""
Converts given messages from OpenAI format to Gemini format
@ -983,7 +984,9 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915
or assistant_msg.get("function_call") is not None
): # support assistant tool invoke conversion
gemini_tool_call_parts = convert_to_gemini_tool_call_invoke(
assistant_msg, model=model
assistant_msg,
model=model,
custom_llm_provider=custom_llm_provider,
)
## check if gemini_tool_call already exists in assistant_content
for gemini_tool_call_part in gemini_tool_call_parts:
@ -1042,7 +1045,10 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915
and messages[msg_i]["role"] in tool_call_message_roles
):
_part = convert_to_gemini_tool_call_result(
messages[msg_i], last_message_with_tool_calls # type: ignore
messages[msg_i], # type: ignore
last_message_with_tool_calls, # type: ignore
model=model,
custom_llm_provider=custom_llm_provider,
)
msg_i += 1
# Handle both single part and list of parts (for Computer Use with images)
@ -1067,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

@ -280,6 +280,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
- gemini-3-pro-preview
- gemini-3-flash
- gemini-3-flash-preview (Gemini 3 Flash)
- gemini-3.1-pro-preview, gemini-3.1-flash, gemini-3.1-flash-lite-preview
- gemini-3.5-flash
- Any future Gemini 3.x models
"""
# Check for Gemini 3 models
@ -287,6 +289,20 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return True
return False
@staticmethod
def _forward_gemini_function_call_id(
model: str, custom_llm_provider: Optional[str] = None
) -> bool:
"""
Whether to include `id` on function_call / function_response parts.
Gemini 3+ on Google AI Studio accepts (and returns) `id` for strict
tool-call matching. Vertex AI rejects the field with HTTP 400.
"""
if custom_llm_provider != "gemini":
return False
return VertexGeminiConfig._is_gemini_3_or_newer(model)
def _supports_penalty_parameters(self, model: str) -> bool:
# Gemini 3 models do not support penalty parameters
if VertexGeminiConfig._is_gemini_3_or_newer(model):
@ -300,6 +316,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
supported_params = [
"temperature",
"top_p",
"top_k",
"max_tokens",
"max_completion_tokens",
"stream",
@ -363,6 +380,66 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
"""
return Tools(googleSearch={})
@staticmethod
def _search_tool_keys() -> set:
return {
VertexToolName.GOOGLE_SEARCH.value,
VertexToolName.GOOGLE_SEARCH_RETRIEVAL.value,
VertexToolName.ENTERPRISE_WEB_SEARCH.value,
VertexToolName.URL_CONTEXT.value,
"google_search",
"google_search_retrieval",
"enterprise_web_search",
"urlContext",
}
@classmethod
def _drop_search_tools_mixed_with_functions(cls, optional_params: dict) -> None:
"""
Drop search tools from optional_params when mixed with function declarations
and include_server_side_tool_invocations is not enabled.
Runs after map_openai_params merges tools and web_search_options so both
code paths (single _map_function call vs split tools + web_search_options)
get the same conflict resolution.
"""
if optional_params.get("include_server_side_tool_invocations"):
return
tools = optional_params.get("tools")
if not isinstance(tools, list) or not tools:
return
search_tool_keys = cls._search_tool_keys()
has_function_declarations = any(
isinstance(tool, dict) and tool.get("function_declarations")
for tool in tools
)
if not has_function_declarations:
return
has_search_tools = any(
isinstance(tool, dict) and any(key in tool for key in search_tool_keys)
for tool in tools
)
if not has_search_tools:
return
verbose_logger.warning(
"Vertex AI does not support mixing function declarations with "
"search tools (googleSearch, enterpriseWebSearch, urlContext, "
"googleSearchRetrieval) in the same request. Dropping search "
"tools and keeping function declarations. To use search tools, "
"send a request without function calling tools."
)
optional_params["tools"] = [
tool
for tool in tools
if not (
isinstance(tool, dict) and any(key in tool for key in search_tool_keys)
)
]
def _map_service_tier_param(self, value: str, optional_params: dict) -> None:
"""
Map OpenAI service_tier (string) to Gemini serviceTier.
@ -884,9 +961,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
GeminiThinkingConfig with thinkingLevel and includeThoughts
"""
# Check if this is gemini-3-flash which supports MINIMAL thinking level
# Covers gemini-3-flash, gemini-3-flash-preview, gemini-3.1-flash, gemini-3.1-flash-lite-preview, etc.
# Covers gemini-3-flash, gemini-3-flash-preview, gemini-3.1-flash, gemini-3.1-flash-lite-preview,
# gemini-3.5-flash, and any future 3.x-flash variants.
is_gemini3flash = model and (
"gemini-3-flash" in model.lower() or "gemini-3.1-flash" in model.lower()
"flash" in model.lower() and "gemini-3" in model.lower()
)
is_gemini31pro = model and ("gemini-3.1-pro-preview" in model.lower())
if reasoning_effort == "minimal":
@ -982,8 +1060,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
# Follow provider defaults unless explicitly opted into legacy behavior.
if litellm.enable_gemini_default_thinking_level_low is True:
is_gemini3flash = (
"gemini-3-flash-preview" in model.lower()
or "gemini-3-flash" in model.lower()
"gemini-3" in model.lower() and "flash" in model.lower()
)
params["thinkingLevel"] = (
"minimal" if is_gemini3flash else "low"
@ -1077,6 +1154,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
model: str,
drop_params: bool,
) -> Dict:
gemini_sampling_params_warned: bool = False
for param, value in non_default_params.items():
if param == "temperature":
if VertexGeminiConfig._is_gemini_3_or_newer(model):
@ -1086,9 +1164,41 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
"can cause infinite loops, degraded reasoning performance, and failure on complex tasks. "
"Strongly recommended to use temperature = 1.0 (default)."
)
if not gemini_sampling_params_warned:
verbose_logger.warning(
"DeprecationWarning: `temperature`, `top_p`, and `top_k` continue to "
f"function for Gemini 3+ ({model}) but are planned for removal in a "
"future release. Move sampling guidance into the `system` "
"instructions instead."
)
gemini_sampling_params_warned = True
optional_params["temperature"] = value
elif param == "top_p":
if (
VertexGeminiConfig._is_gemini_3_or_newer(model)
and not gemini_sampling_params_warned
):
verbose_logger.warning(
"DeprecationWarning: `temperature`, `top_p`, and `top_k` continue to "
f"function for Gemini 3+ ({model}) but are planned for removal in a "
"future release. Move sampling guidance into the `system` "
"instructions instead."
)
gemini_sampling_params_warned = True
optional_params["top_p"] = value
elif param == "top_k":
if (
VertexGeminiConfig._is_gemini_3_or_newer(model)
and not gemini_sampling_params_warned
):
verbose_logger.warning(
"DeprecationWarning: `temperature`, `top_p`, and `top_k` continue to "
f"function for Gemini 3+ ({model}) but are planned for removal in a "
"future release. Move sampling guidance into the `system` "
"instructions instead."
)
gemini_sampling_params_warned = True
optional_params["top_k"] = value
elif (
param == "stream" and value is True
): # sending stream = False, can cause it to get passed unchecked and raise issues
@ -1139,11 +1249,14 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
if _tool_choice_value is not None:
optional_params["tool_choice"] = _tool_choice_value
elif param == "parallel_tool_calls":
if value is False and not (
drop_params or litellm.drop_params
): # if drop params is True, then we should just ignore this
self.validate_parallel_tool_calls(value, non_default_params)
else:
tools_list = non_default_params.get(
"tools", non_default_params.get("functions")
)
num_tools = len(tools_list) if isinstance(tools_list, list) else 0
# Gemini does not support parallel_tool_calls=False with multiple
# tools. Drop the param instead of failing — Responses API clients
# often send parallel_tool_calls=false by default.
if not (value is False and num_tools > 1):
optional_params["parallel_tool_calls"] = value
elif param == "seed":
optional_params["seed"] = value
@ -1216,6 +1329,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
if "temperature" not in optional_params:
optional_params["temperature"] = 1.0
self._drop_search_tools_mixed_with_functions(optional_params)
return optional_params
def get_mapped_special_auth_params(self) -> dict:
@ -1588,6 +1703,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
}
# Extract thought signature if present
thought_signature = part.get("thoughtSignature")
# Gemini 3.5+ returns a stable `id` per function call to enable
# strict response matching. Preserve it as the OpenAI
# tool_call_id so it can be echoed back unchanged.
gemini_call_id = part["functionCall"].get("id")
if is_function_call is True:
function_dict: Dict[str, Any] = dict(_function_chunk)
@ -1605,6 +1724,11 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
"function": _function_chunk,
"index": cumulative_tool_call_idx,
}
# Gemini 3.5+ returns a stable native `id`; prefer it over
# the synthetic call_<uuid> so the same value can be echoed
# back on the matching `functionResponse`.
if gemini_call_id:
_tool_response_chunk["id"] = gemini_call_id
# Embed thought signature in ID for OpenAI client compatibility
if thought_signature:
_tool_response_chunk["provider_specific_fields"] = { # type: ignore
@ -2539,7 +2663,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
litellm_params: Optional[dict] = None,
) -> List[ContentType]:
return _gemini_convert_messages_with_history(
messages=messages, model=model, litellm_params=litellm_params
messages=messages,
model=model,
litellm_params=litellm_params,
custom_llm_provider="vertex_ai",
)
def get_error_class(

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

@ -45,7 +45,7 @@ class PartnerModelPrefixes(str, Enum):
class VertexAIPartnerModels(VertexBase):
def __init__(self) -> None:
pass
super().__init__()
@staticmethod
def is_vertex_partner_model(model: str):
@ -116,9 +116,6 @@ class VertexAIPartnerModels(VertexBase):
CodestralTextCompletion,
)
from litellm.llms.openai_like.chat.handler import OpenAILikeChatHandler
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexLLM,
)
except Exception as e:
raise VertexAIError(
status_code=400,
@ -133,9 +130,7 @@ class VertexAIPartnerModels(VertexBase):
message="""Upgrade vertex ai. Run `pip install "google-cloud-aiplatform>=1.38"`""",
)
try:
vertex_httpx_logic = VertexLLM()
access_token, project_id = vertex_httpx_logic._ensure_access_token(
access_token, project_id = self._ensure_access_token(
credentials=vertex_credentials,
project_id=vertex_project,
custom_llm_provider="vertex_ai",

View file

@ -31,7 +31,7 @@ from ..vertex_llm_base import VertexBase
class VertexAIGemmaModels(VertexBase):
def __init__(self) -> None:
pass
super().__init__()
def completion(
self,
@ -62,9 +62,6 @@ class VertexAIGemmaModels(VertexBase):
try:
import vertexai
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexLLM,
)
from litellm.llms.vertex_ai.vertex_gemma_models.transformation import (
VertexGemmaConfig,
)
@ -83,9 +80,8 @@ class VertexAIGemmaModels(VertexBase):
)
try:
model = get_vertex_base_model_name(model=model)
vertex_httpx_logic = VertexLLM()
access_token, project_id = vertex_httpx_logic._ensure_access_token(
access_token, project_id = self._ensure_access_token(
credentials=vertex_credentials,
project_id=vertex_project,
custom_llm_provider="vertex_ai",

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