mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-24 00:52:24 +00:00
Merge branch 'litellm_internal_staging' into litellm_batch_with_policy
This commit is contained in:
commit
b6dd5c3ce9
744 changed files with 36983 additions and 5045 deletions
|
|
@ -2477,10 +2477,15 @@ jobs:
|
|||
DISABLE_SCHEMA_UPDATE: "true"
|
||||
SERVER_ROOT_PATH: ""
|
||||
PROXY_LOGOUT_URL: ""
|
||||
# LITELLM_LICENSE is forwarded from the project env so premium-gated
|
||||
# UI flows can be exercised. license.spec.ts asserts the resulting
|
||||
# JWT carries premium_user=true; if it ever stops being passed, that
|
||||
# test fails loudly rather than silently regressing premium coverage.
|
||||
command: |
|
||||
uv run --no-sync python -m litellm.proxy.proxy_cli \
|
||||
--config ui/litellm-dashboard/e2e_tests/fixtures/config.yml \
|
||||
--port 4000
|
||||
LITELLM_LICENSE="$LITELLM_LICENSE" \
|
||||
uv run --no-sync python -m litellm.proxy.proxy_cli \
|
||||
--config ui/litellm-dashboard/e2e_tests/fixtures/config.yml \
|
||||
--port 4000
|
||||
background: true
|
||||
- run:
|
||||
name: Wait for proxy to be ready
|
||||
|
|
@ -2497,9 +2502,12 @@ jobs:
|
|||
exit 1
|
||||
- run:
|
||||
name: Run Playwright E2E tests
|
||||
# Forward LITELLM_LICENSE so license.spec.ts can detect that the
|
||||
# proxy was launched with a license and assert premium_user=true.
|
||||
command: |
|
||||
cd ui/litellm-dashboard
|
||||
npx playwright test --config e2e_tests/playwright.config.ts
|
||||
LITELLM_LICENSE="$LITELLM_LICENSE" \
|
||||
npx playwright test --config e2e_tests/playwright.config.ts
|
||||
no_output_timeout: 10m
|
||||
- store_artifacts:
|
||||
path: ui/litellm-dashboard/test-results
|
||||
|
|
@ -2533,7 +2541,6 @@ jobs:
|
|||
paths:
|
||||
- litellm-docker-database.tar.zst
|
||||
|
||||
|
||||
test_bad_database_url:
|
||||
machine:
|
||||
image: ubuntu-2204:2024.04.1
|
||||
|
|
|
|||
28
.github/workflows/codeql.yml
vendored
28
.github/workflows/codeql.yml
vendored
|
|
@ -53,3 +53,31 @@ jobs:
|
|||
uses: github/codeql-action/analyze@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3
|
||||
with:
|
||||
category: "/language:${{ matrix.language }}"
|
||||
output: sarif-results
|
||||
upload: failure-only
|
||||
|
||||
# py/weak-sensitive-data-hashing (CWE-328) fires on the OCI signing call at
|
||||
# litellm/llms/oci/common_utils.py, which hashes the HTTP request body to
|
||||
# produce the x-content-sha256 header required by the OCI HTTP signing spec —
|
||||
# a content-integrity hash, not a password or secret hash. SHA-256 is mandated
|
||||
# by Oracle for this header; see
|
||||
# https://docs.oracle.com/en-us/iaas/Content/API/Concepts/signingrequests.htm
|
||||
# The `usedforsecurity=False` flag on the hashlib.sha256 call already declares
|
||||
# non-security intent, but CodeQL's taint flow still re-fires when callers
|
||||
# further up the stack are modified. The suppression is scoped to this one
|
||||
# file/rule pair via SARIF post-filtering so every other callsite of
|
||||
# py/weak-sensitive-data-hashing in the repository continues to be analyzed.
|
||||
- name: Filter SARIF (OCI sha256)
|
||||
if: matrix.language == 'python'
|
||||
uses: advanced-security/filter-sarif@2da736ff05ef065cb2894ac6892e47b5eac2c3c0 # v1.1
|
||||
with:
|
||||
patterns: |
|
||||
-litellm/llms/oci/common_utils.py:py/weak-sensitive-data-hashing
|
||||
input: sarif-results/python.sarif
|
||||
output: sarif-results/python.sarif
|
||||
|
||||
- name: Upload SARIF
|
||||
uses: github/codeql-action/upload-sarif@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3
|
||||
with:
|
||||
sarif_file: sarif-results
|
||||
category: "/language:${{ matrix.language }}"
|
||||
|
|
|
|||
1
.github/workflows/test-unit-proxy-db.yml
vendored
1
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -218,6 +218,7 @@ jobs:
|
|||
tests/proxy_unit_tests/test_gemini_agents_endpoints.py
|
||||
tests/proxy_unit_tests/test_get_favicon.py
|
||||
tests/proxy_unit_tests/test_get_image.py
|
||||
tests/proxy_unit_tests/test_reducto_ocr_route.py
|
||||
tests/proxy_unit_tests/test_ui_path_detection.py
|
||||
tests/proxy_unit_tests/test_prompt_test_endpoint.py
|
||||
tests/proxy_unit_tests/test_check_batch_cost.py
|
||||
|
|
|
|||
34
.github/workflows/test-unit-proxy-mgmt-behavior.yml
vendored
Normal file
34
.github/workflows/test-unit-proxy-mgmt-behavior.yml
vendored
Normal file
|
|
@ -0,0 +1,34 @@
|
|||
name: "Unit Tests: Proxy Management-Endpoint Behavior Pinning"
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_branch
|
||||
- "litellm_**"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
proxy-mgmt-behavior:
|
||||
uses: ./.github/workflows/_test-unit-services-base.yml
|
||||
with:
|
||||
test-path: tests/proxy_behavior
|
||||
# workers=0 (no xdist): the world seed is a single shared Postgres
|
||||
# state — two xdist workers both call seed_world() and race on the
|
||||
# ``behavior-pin-budget`` row, producing UniqueViolation + cascading
|
||||
# missing-membership FK failures. The whole suite is ~7s sequentially,
|
||||
# so the cost of disabling parallelism here is negligible.
|
||||
workers: 0
|
||||
reruns: 0
|
||||
enable-postgres: true
|
||||
artifact-name: proxy-mgmt-behavior
|
||||
timeout-minutes: 15
|
||||
|
|
@ -24,7 +24,8 @@ RUN for i in 1 2 3; do \
|
|||
curl \
|
||||
openssl \
|
||||
libsndfile \
|
||||
nodejs && break || sleep 5; \
|
||||
nodejs \
|
||||
npm && break || sleep 5; \
|
||||
done
|
||||
|
||||
ENV UV_PROJECT_ENVIRONMENT=/app/.venv \
|
||||
|
|
|
|||
File diff suppressed because one or more lines are too long
|
|
@ -413,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
|
||||
|
|
@ -636,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()
|
||||
|
||||
|
||||
|
|
@ -903,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)
|
||||
|
||||
|
|
@ -1014,6 +1023,7 @@ model_list = list(
|
|||
| ovhcloud_models
|
||||
| lemonade_models
|
||||
| docker_model_runner_models
|
||||
| reducto_models
|
||||
| bedrock_mantle_models
|
||||
| set(clarifai_models)
|
||||
)
|
||||
|
|
@ -1120,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,
|
||||
}
|
||||
|
||||
|
|
@ -1866,6 +1877,9 @@ if TYPE_CHECKING:
|
|||
from .llms.azure.completion.transformation import (
|
||||
AzureOpenAITextConfig as AzureOpenAITextConfig,
|
||||
)
|
||||
from .llms.azure.audio_transcription.transformation import (
|
||||
AzureSpeechAudioTranscriptionConfig as AzureSpeechAudioTranscriptionConfig,
|
||||
)
|
||||
from .llms.hosted_vllm.chat.transformation import (
|
||||
HostedVLLMChatConfig as HostedVLLMChatConfig,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -273,6 +273,7 @@ LLM_CONFIG_NAMES = (
|
|||
"AzureOpenAIConfig",
|
||||
"AzureOpenAIGPT5Config",
|
||||
"AzureOpenAITextConfig",
|
||||
"AzureSpeechAudioTranscriptionConfig",
|
||||
"HostedVLLMChatConfig",
|
||||
"HostedVLLMEmbeddingConfig",
|
||||
# Alias for backwards compatibility
|
||||
|
|
@ -1054,6 +1055,10 @@ _LLM_CONFIGS_IMPORT_MAP = {
|
|||
".llms.azure.completion.transformation",
|
||||
"AzureOpenAITextConfig",
|
||||
),
|
||||
"AzureSpeechAudioTranscriptionConfig": (
|
||||
".llms.azure.audio_transcription.transformation",
|
||||
"AzureSpeechAudioTranscriptionConfig",
|
||||
),
|
||||
"HostedVLLMChatConfig": (
|
||||
".llms.hosted_vllm.chat.transformation",
|
||||
"HostedVLLMChatConfig",
|
||||
|
|
|
|||
|
|
@ -6,7 +6,6 @@ Always uses fastuuid for performance.
|
|||
|
||||
import fastuuid as _uuid # type: ignore
|
||||
|
||||
|
||||
# Expose a module-like alias so callers can use: uuid.uuid4()
|
||||
uuid = _uuid
|
||||
|
||||
|
|
|
|||
|
|
@ -9,7 +9,6 @@ from typing import Dict, Optional
|
|||
|
||||
from .exceptions import AnthropicErrorResponse, AnthropicErrorType
|
||||
|
||||
|
||||
# HTTP status code -> Anthropic error type
|
||||
# Source: https://docs.anthropic.com/en/api/errors
|
||||
ANTHROPIC_ERROR_TYPE_MAP: Dict[int, AnthropicErrorType] = {
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@
|
|||
|
||||
from typing_extensions import Literal, Required, TypedDict
|
||||
|
||||
|
||||
# Known Anthropic error types
|
||||
# Source: https://docs.anthropic.com/en/api/errors
|
||||
AnthropicErrorType = Literal[
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -5,7 +5,6 @@ Auto-detect content type per message: code, JSON, or text.
|
|||
import json
|
||||
import re
|
||||
|
||||
|
||||
_CODE_KEYWORDS = re.compile(
|
||||
r"\b(?:def |function |class |import |from |require\(|#include|fn |func |const |let |var |public |private |static )\b"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1443,6 +1443,12 @@ CLI_JWT_EXPIRATION_HOURS = int(
|
|||
or os.getenv("LITELLM_CLI_JWT_EXPIRATION_HOURS")
|
||||
or 24
|
||||
)
|
||||
# Comma-separated allowlisted OIDC claim map for CLI SSO polling, e.g.
|
||||
# "employment_type->acme_employment_type,org_info.department->department"
|
||||
CLI_SSO_CLAIM_MAP = (
|
||||
os.getenv("CLI_SSO_CLAIM_MAP") or os.getenv("LITELLM_CLI_SSO_CLAIM_MAP") or ""
|
||||
)
|
||||
CLI_SSO_CLAIM_MAX_SCALAR_LENGTH = 1024
|
||||
|
||||
########################### UI SESSION DURATION ###########################
|
||||
# Duration for UI login session (username/password, SSO, invitation links). Format: "30s", "30m", "24h", "7d"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
from typing import AsyncIterator, Dict, Iterator, Literal, NamedTuple, Union
|
||||
|
||||
|
||||
FileContentProvider = Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "anthropic", "manus"
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,10 +1,10 @@
|
|||
"""
|
||||
Google GenAI Adapters for LiteLLM
|
||||
|
||||
This module provides adapters for transforming Google GenAI generate_content requests
|
||||
This module provides adapters for transforming Google GenAI generate_content requests
|
||||
to/from LiteLLM completion format with full support for:
|
||||
- Text content transformation
|
||||
- Tool calling (function declarations, function calls, function responses)
|
||||
- Tool calling (function declarations, function calls, function responses)
|
||||
- Streaming (both regular and tool calling)
|
||||
- Mixed content (text + tool calls)
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
"""
|
||||
Handles Batching + sending Httpx Post requests to slack
|
||||
Handles Batching + sending Httpx Post requests to slack
|
||||
|
||||
Slack alerts are sent every 10s or when events are greater than X events
|
||||
Slack alerts are sent every 10s or when events are greater than X events
|
||||
|
||||
see custom_batch_logger.py for more details / defaults
|
||||
see custom_batch_logger.py for more details / defaults
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ else:
|
|||
|
||||
|
||||
def process_slack_alerting_variables(
|
||||
alert_to_webhook_url: Optional[Dict[AlertType, Union[List[str], str]]]
|
||||
alert_to_webhook_url: Optional[Dict[AlertType, Union[List[str], str]]],
|
||||
) -> Optional[Dict[AlertType, Union[List[str], str]]]:
|
||||
"""
|
||||
process alert_to_webhook_url
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Base class for Additional Logging Utils for CustomLoggers
|
||||
Base class for Additional Logging Utils for CustomLoggers
|
||||
|
||||
- Health Check for the logging util
|
||||
- Get Request / Response Payload for the logging util
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Custom Logger that handles batching logic
|
||||
Custom Logger that handles batching logic
|
||||
|
||||
Use this if you want your logs to be stored in memory and flushed periodically.
|
||||
"""
|
||||
|
|
@ -14,22 +14,38 @@ from litellm.integrations.custom_logger import CustomLogger
|
|||
|
||||
|
||||
class CustomBatchLogger(CustomLogger):
|
||||
preserve_events_added_during_flush = False
|
||||
|
||||
# Default cap on the in-memory log queue. Prevents unbounded memory growth
|
||||
# if ``async_send_batch`` consistently fails (e.g. the destination is
|
||||
# unreachable) and events are preserved across flush attempts. Subclasses
|
||||
# may override by passing ``max_queue_size`` or by setting the attribute
|
||||
# directly (see ``RubrikLogger`` for an example).
|
||||
DEFAULT_MAX_QUEUE_SIZE = 50_000
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
flush_lock: Optional[asyncio.Lock] = None,
|
||||
batch_size: Optional[int] = None,
|
||||
flush_interval: Optional[int] = None,
|
||||
max_queue_size: Optional[int] = None,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
"""
|
||||
Args:
|
||||
flush_lock (Optional[asyncio.Lock], optional): Lock to use when flushing the queue. Defaults to None. Only used for custom loggers that do batching
|
||||
max_queue_size (Optional[int], optional): Maximum number of events to retain in ``log_queue``. When the limit is exceeded (e.g. because the send destination is unreachable and events are preserved for retry), the oldest events are dropped. Defaults to ``DEFAULT_MAX_QUEUE_SIZE``.
|
||||
"""
|
||||
self.log_queue: List = []
|
||||
self.flush_interval = flush_interval or litellm.DEFAULT_FLUSH_INTERVAL_SECONDS
|
||||
self.batch_size: int = batch_size or litellm.DEFAULT_BATCH_SIZE
|
||||
self.last_flush_time = time.time()
|
||||
self.flush_lock = flush_lock
|
||||
self.max_queue_size: int = (
|
||||
max_queue_size
|
||||
if max_queue_size is not None
|
||||
else self.DEFAULT_MAX_QUEUE_SIZE
|
||||
)
|
||||
|
||||
super().__init__(**kwargs)
|
||||
|
||||
|
|
@ -47,11 +63,40 @@ class CustomBatchLogger(CustomLogger):
|
|||
|
||||
async with self.flush_lock:
|
||||
if self.log_queue:
|
||||
log_queue_length = len(self.log_queue)
|
||||
verbose_logger.debug(
|
||||
"CustomLogger: Flushing batch of %s events", len(self.log_queue)
|
||||
)
|
||||
await self.async_send_batch()
|
||||
self.log_queue.clear()
|
||||
try:
|
||||
await self.async_send_batch()
|
||||
except Exception:
|
||||
# If the underlying batch send raised, do NOT drop the
|
||||
# in-flight events. They will be retried on the next flush.
|
||||
# Most existing async_send_batch implementations swallow
|
||||
# their own errors, so this only affects loggers that opt
|
||||
# in to surfacing failures (e.g. Rubrik).
|
||||
verbose_logger.exception(
|
||||
"CustomLogger: async_send_batch raised; preserving "
|
||||
"%s events in queue for retry",
|
||||
log_queue_length,
|
||||
)
|
||||
# Guard against unbounded queue growth if the destination
|
||||
# is persistently unreachable. Drop the oldest events
|
||||
# beyond ``max_queue_size``.
|
||||
overflow = len(self.log_queue) - self.max_queue_size
|
||||
if overflow > 0:
|
||||
del self.log_queue[:overflow]
|
||||
verbose_logger.warning(
|
||||
"CustomLogger: log queue exceeded max_queue_size=%s; "
|
||||
"dropped %s oldest events.",
|
||||
self.max_queue_size,
|
||||
overflow,
|
||||
)
|
||||
return
|
||||
if self.preserve_events_added_during_flush:
|
||||
del self.log_queue[:log_queue_length]
|
||||
else:
|
||||
self.log_queue.clear()
|
||||
self.last_flush_time = time.time()
|
||||
|
||||
async def async_send_batch(self, *args, **kwargs):
|
||||
|
|
|
|||
|
|
@ -9,7 +9,6 @@ import polars as pl
|
|||
|
||||
from .schema import FOCUS_NORMALIZED_SCHEMA
|
||||
|
||||
|
||||
_TAG_KEYS = (
|
||||
"team_id",
|
||||
"team_alias",
|
||||
|
|
|
|||
|
|
@ -702,6 +702,14 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
},
|
||||
)
|
||||
|
||||
# _record_exception_on_span only stamps when error_code is set;
|
||||
# bare TypeError etc. has none, and the span is about to be ended.
|
||||
error_code = (
|
||||
error_information.get("error_code") if error_information else None
|
||||
)
|
||||
if not error_code:
|
||||
self.set_response_status_code_attribute(parent_otel_span, 500)
|
||||
|
||||
# Pre-request latency (request_data carries the propagated
|
||||
# metadata on the failure path; omitted if it failed before handoff).
|
||||
self.set_preprocessing_duration_attribute(parent_otel_span, request_data)
|
||||
|
|
@ -726,9 +734,57 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
exception_logging_span.set_status(Status(StatusCode.ERROR))
|
||||
exception_logging_span.end(end_time=self._to_ns(datetime.now()))
|
||||
|
||||
# Emit guardrail spans for any guardrail invocations that
|
||||
# ran during this request. _handle_failure typically does this,
|
||||
# but for pre-call guardrail blocks the standard_logging_object
|
||||
# may not carry guardrail_information by the time _handle_failure
|
||||
# fires (the data lives only in request_data["metadata"]). Pull
|
||||
# directly from request_data so the span is recorded either way;
|
||||
# _emit_once dedupes if _handle_failure already emitted it.
|
||||
self._emit_guardrail_spans_from_request_data(
|
||||
request_data=request_data,
|
||||
parent_span=parent_otel_span,
|
||||
)
|
||||
|
||||
# End Parent OTEL Sspan
|
||||
parent_otel_span.end(end_time=self._to_ns(datetime.now()))
|
||||
|
||||
def _emit_guardrail_spans_from_request_data(
|
||||
self,
|
||||
request_data: dict,
|
||||
parent_span: Optional[Any],
|
||||
) -> None:
|
||||
"""Emit ``guardrail`` spans from ``request_data["metadata"]
|
||||
["standard_logging_guardrail_information"]``.
|
||||
|
||||
Routed through ``_create_guardrail_span`` so the dedupe state in
|
||||
``_otel_internal`` is honoured — if ``_handle_failure`` already
|
||||
emitted these spans for the same kwargs, this is a no-op.
|
||||
"""
|
||||
from opentelemetry import trace as _trace
|
||||
|
||||
metadata = (request_data or {}).get("metadata") or {}
|
||||
guardrail_information = metadata.get("standard_logging_guardrail_information")
|
||||
if not guardrail_information:
|
||||
return
|
||||
|
||||
# _create_guardrail_span reads guardrail_information from
|
||||
# kwargs["standard_logging_object"] and shares its dedupe state via
|
||||
# kwargs["litellm_params"]["metadata"]["_otel_internal"]. Pass the
|
||||
# SAME metadata dict the proxy populated so _handle_failure and
|
||||
# this hook see the same dedupe markers.
|
||||
kwargs: Dict[str, Any] = {
|
||||
"litellm_params": {"metadata": metadata},
|
||||
"standard_logging_object": {
|
||||
"guardrail_information": guardrail_information,
|
||||
"metadata": metadata,
|
||||
},
|
||||
}
|
||||
context = (
|
||||
_trace.set_span_in_context(parent_span) if parent_span is not None else None
|
||||
)
|
||||
self._create_guardrail_span(kwargs=kwargs, context=context)
|
||||
|
||||
async def async_post_call_success_hook(
|
||||
self,
|
||||
data: dict,
|
||||
|
|
@ -750,11 +806,6 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
# Pre-request latency on the SERVER span (success path).
|
||||
self.set_preprocessing_duration_attribute(parent_span, kwargs)
|
||||
|
||||
# http.response.status_code on the SERVER span (success path).
|
||||
# A successful proxy response is HTTP 200; the failure path sets
|
||||
# this from the error code in _record_exception_on_span.
|
||||
self.set_response_status_code_attribute(parent_span, 200)
|
||||
|
||||
# 3. Guardrail span
|
||||
self._create_guardrail_span(kwargs=kwargs, context=ctx)
|
||||
|
||||
|
|
@ -937,7 +988,15 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
and hasattr(proxy_span, "is_recording")
|
||||
and proxy_span.is_recording()
|
||||
):
|
||||
proxy_span.end(end_time=self._to_ns(end_time))
|
||||
self._close_proxy_span_ok(proxy_span, end_time)
|
||||
|
||||
def _close_proxy_span_ok(self, span: Span, end_time) -> None:
|
||||
"""Stamp http.response.status_code=200 + status=OK, then end the span."""
|
||||
from opentelemetry.trace import Status, StatusCode
|
||||
|
||||
self.set_response_status_code_attribute(span, 200)
|
||||
span.set_status(Status(StatusCode.OK))
|
||||
span.end(end_time=self._to_ns(end_time))
|
||||
|
||||
def _handle_success(self, kwargs, response_obj, start_time, end_time):
|
||||
"""Create the litellm_request span then close the proxy span."""
|
||||
|
|
@ -1023,8 +1082,10 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
parent_span is not None
|
||||
and hasattr(parent_span, "name")
|
||||
and parent_span.name == LITELLM_PROXY_REQUEST_SPAN_NAME
|
||||
and hasattr(parent_span, "is_recording")
|
||||
and parent_span.is_recording()
|
||||
):
|
||||
parent_span.end(end_time=self._to_ns(end_time))
|
||||
self._close_proxy_span_ok(parent_span, end_time)
|
||||
|
||||
# Stamp team attributes onto the SERVER (root) span before it is
|
||||
# closed, so the trace root carries them like every child span.
|
||||
|
|
@ -1617,6 +1678,37 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
"guardrail_response", safe_dumps(guardrail_response)
|
||||
)
|
||||
|
||||
# Surface guardrail_status (success / guardrail_intervened /
|
||||
# guardrail_failed_to_respond / not_run) as a top-level span
|
||||
# attribute so trace backends can filter on it without parsing
|
||||
# guardrail_response.
|
||||
self.safe_set_attribute(
|
||||
span=guardrail_span,
|
||||
key="guardrail_status",
|
||||
value=guardrail_information.get("guardrail_status"),
|
||||
)
|
||||
|
||||
# Provider's raw top-level action (e.g. Bedrock's
|
||||
# ``GUARDRAIL_INTERVENED`` / ``NONE``). Populated by the provider
|
||||
# hook onto StandardLoggingGuardrailInformation so this integration
|
||||
# stays provider-agnostic — we only read a normalised string.
|
||||
guardrail_action = guardrail_information.get("guardrail_action")
|
||||
if guardrail_action:
|
||||
guardrail_span.set_attribute("guardrail_action", guardrail_action)
|
||||
|
||||
# The provider hook (e.g. Bedrock) extracts violation_categories
|
||||
# from the raw response BEFORE redaction and stamps them onto
|
||||
# StandardLoggingGuardrailInformation. Surfacing them here as a
|
||||
# queryable attribute lets dashboards group by violation category
|
||||
# without parsing the redacted guardrail_response blob.
|
||||
violation_categories = guardrail_information.get("violation_categories")
|
||||
if violation_categories:
|
||||
# OTel sequence attributes must be homogeneous primitives;
|
||||
# serialise to JSON once so set_attribute never coerces.
|
||||
guardrail_span.set_attribute(
|
||||
"guardrail_violation_categories", safe_dumps(violation_categories)
|
||||
)
|
||||
|
||||
self._set_team_attributes_from_kwargs(guardrail_span, kwargs)
|
||||
|
||||
guardrail_span.end(end_time=self._to_ns(end_time_datetime))
|
||||
|
|
@ -2962,6 +3054,11 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
management_endpoint_span.set_status(Status(StatusCode.OK))
|
||||
management_endpoint_span.end(end_time=_end_time_ns)
|
||||
|
||||
# The management wrapper has no other hook that closes the SERVER span.
|
||||
self.set_response_status_code_attribute(parent_otel_span, 200)
|
||||
parent_otel_span.set_status(Status(StatusCode.OK))
|
||||
parent_otel_span.end(end_time=_end_time_ns)
|
||||
|
||||
async def async_management_endpoint_failure_hook(
|
||||
self,
|
||||
logging_payload: ManagementEndpointLoggingPayload,
|
||||
|
|
@ -3012,6 +3109,24 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
|
|||
management_endpoint_span.set_status(Status(StatusCode.ERROR))
|
||||
management_endpoint_span.end(end_time=_end_time_ns)
|
||||
|
||||
# The management wrapper has no other hook that closes the SERVER span.
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
StandardLoggingPayloadSetup,
|
||||
)
|
||||
|
||||
error_information = StandardLoggingPayloadSetup.get_error_information(
|
||||
original_exception=_exception,
|
||||
)
|
||||
parent_otel_span.set_status(Status(StatusCode.ERROR))
|
||||
self._record_exception_on_span(
|
||||
span=parent_otel_span,
|
||||
kwargs={
|
||||
"exception": _exception,
|
||||
"standard_logging_object": {"error_information": error_information},
|
||||
},
|
||||
)
|
||||
parent_otel_span.end(end_time=_end_time_ns)
|
||||
|
||||
def create_litellm_proxy_request_started_span(
|
||||
self,
|
||||
start_time: datetime,
|
||||
|
|
|
|||
|
|
@ -105,7 +105,7 @@ def _remove_nulls(x: Dict[str, Any]) -> Dict[str, Any]:
|
|||
|
||||
|
||||
def get_traces_and_spans_from_payload(
|
||||
payload: List[Dict[str, Any]]
|
||||
payload: List[Dict[str, Any]],
|
||||
) -> Tuple[List[Dict[str, Any]], List[Dict[str, Any]]]:
|
||||
"""
|
||||
Separate traces and spans from payload.
|
||||
|
|
|
|||
|
|
@ -166,6 +166,53 @@ class PrometheusLogger(CustomLogger):
|
|||
labelnames=self.get_labels_for_metric("litellm_output_tokens_metric"),
|
||||
)
|
||||
|
||||
# Token-type detail metrics. These break out cached, cache-creation,
|
||||
# audio and reasoning tokens that providers report inside
|
||||
# prompt_tokens_details / completion_tokens_details on the usage
|
||||
# object. They are sparse (only incremented when the provider
|
||||
# reports a non-zero value) and are additive to the existing
|
||||
# input/output token totals — no breaking change for existing
|
||||
# dashboards built on the totals.
|
||||
self.litellm_input_cached_tokens_metric = self._counter_factory(
|
||||
"litellm_input_cached_tokens_metric",
|
||||
"Provider-side cached input tokens (e.g. OpenAI prompt_tokens_details.cached_tokens, Anthropic cache_read_input_tokens)",
|
||||
labelnames=self.get_labels_for_metric(
|
||||
"litellm_input_cached_tokens_metric"
|
||||
),
|
||||
)
|
||||
|
||||
self.litellm_input_cache_creation_tokens_metric = self._counter_factory(
|
||||
"litellm_input_cache_creation_tokens_metric",
|
||||
"Provider-side input tokens written to prompt cache (e.g. Anthropic cache_creation_input_tokens)",
|
||||
labelnames=self.get_labels_for_metric(
|
||||
"litellm_input_cache_creation_tokens_metric"
|
||||
),
|
||||
)
|
||||
|
||||
self.litellm_input_audio_tokens_metric = self._counter_factory(
|
||||
"litellm_input_audio_tokens_metric",
|
||||
"Audio input tokens reported in prompt_tokens_details.audio_tokens",
|
||||
labelnames=self.get_labels_for_metric(
|
||||
"litellm_input_audio_tokens_metric"
|
||||
),
|
||||
)
|
||||
|
||||
self.litellm_output_reasoning_tokens_metric = self._counter_factory(
|
||||
"litellm_output_reasoning_tokens_metric",
|
||||
"Reasoning tokens reported in completion_tokens_details.reasoning_tokens",
|
||||
labelnames=self.get_labels_for_metric(
|
||||
"litellm_output_reasoning_tokens_metric"
|
||||
),
|
||||
)
|
||||
|
||||
self.litellm_output_audio_tokens_metric = self._counter_factory(
|
||||
"litellm_output_audio_tokens_metric",
|
||||
"Audio output tokens reported in completion_tokens_details.audio_tokens",
|
||||
labelnames=self.get_labels_for_metric(
|
||||
"litellm_output_audio_tokens_metric"
|
||||
),
|
||||
)
|
||||
|
||||
# Remaining Budget for Team
|
||||
self.litellm_remaining_team_budget_metric = self._gauge_factory(
|
||||
"litellm_remaining_team_budget_metric",
|
||||
|
|
@ -1301,6 +1348,101 @@ class PrometheusLogger(CustomLogger):
|
|||
amount=float(standard_logging_payload["completion_tokens"]),
|
||||
)
|
||||
|
||||
# Token-type detail metrics — sparse, only emitted when the provider
|
||||
# reports a non-zero value in usage.prompt_tokens_details /
|
||||
# usage.completion_tokens_details.
|
||||
self._increment_token_detail_metrics(
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
enum_values=enum_values,
|
||||
label_context=label_context,
|
||||
)
|
||||
|
||||
def _increment_token_detail_metrics(
|
||||
self,
|
||||
standard_logging_payload: StandardLoggingPayload,
|
||||
enum_values: UserAPIKeyLabelValues,
|
||||
label_context: Optional[PrometheusLabelFactoryContext] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Increment per-token-type counters from the Usage object that providers
|
||||
attach to the request. The Usage dict is plumbed onto
|
||||
``standard_logging_payload["metadata"]["usage_object"]`` by
|
||||
``get_standard_logging_object_payload``.
|
||||
|
||||
Each counter is only incremented when the underlying value is > 0, so
|
||||
scrape output stays sparse for providers that don't report these
|
||||
details (most non-OpenAI/Anthropic models).
|
||||
"""
|
||||
metadata = standard_logging_payload.get("metadata") or {}
|
||||
usage_object = (
|
||||
metadata.get("usage_object") if isinstance(metadata, dict) else None
|
||||
)
|
||||
if not isinstance(usage_object, dict):
|
||||
return
|
||||
|
||||
prompt_details = usage_object.get("prompt_tokens_details") or {}
|
||||
completion_details = usage_object.get("completion_tokens_details") or {}
|
||||
|
||||
detail_metrics: List[Tuple[Any, DEFINED_PROMETHEUS_METRICS, Any]] = [
|
||||
(
|
||||
self.litellm_input_cached_tokens_metric,
|
||||
"litellm_input_cached_tokens_metric",
|
||||
(
|
||||
prompt_details.get("cached_tokens")
|
||||
if isinstance(prompt_details, dict)
|
||||
else None
|
||||
),
|
||||
),
|
||||
(
|
||||
self.litellm_input_cache_creation_tokens_metric,
|
||||
"litellm_input_cache_creation_tokens_metric",
|
||||
(
|
||||
prompt_details.get("cache_creation_tokens")
|
||||
if isinstance(prompt_details, dict)
|
||||
else None
|
||||
),
|
||||
),
|
||||
(
|
||||
self.litellm_input_audio_tokens_metric,
|
||||
"litellm_input_audio_tokens_metric",
|
||||
(
|
||||
prompt_details.get("audio_tokens")
|
||||
if isinstance(prompt_details, dict)
|
||||
else None
|
||||
),
|
||||
),
|
||||
(
|
||||
self.litellm_output_reasoning_tokens_metric,
|
||||
"litellm_output_reasoning_tokens_metric",
|
||||
(
|
||||
completion_details.get("reasoning_tokens")
|
||||
if isinstance(completion_details, dict)
|
||||
else None
|
||||
),
|
||||
),
|
||||
(
|
||||
self.litellm_output_audio_tokens_metric,
|
||||
"litellm_output_audio_tokens_metric",
|
||||
(
|
||||
completion_details.get("audio_tokens")
|
||||
if isinstance(completion_details, dict)
|
||||
else None
|
||||
),
|
||||
),
|
||||
]
|
||||
|
||||
for counter, metric_name, value in detail_metrics:
|
||||
if not isinstance(value, (int, float)) or value <= 0:
|
||||
continue
|
||||
PrometheusLogger._inc_labeled_counter(
|
||||
self,
|
||||
counter,
|
||||
metric_name,
|
||||
enum_values,
|
||||
label_context=label_context,
|
||||
amount=float(value),
|
||||
)
|
||||
|
||||
def _increment_cache_metrics(
|
||||
self,
|
||||
standard_logging_payload: StandardLoggingPayload,
|
||||
|
|
|
|||
605
litellm/integrations/rubrik.py
Normal file
605
litellm/integrations/rubrik.py
Normal file
|
|
@ -0,0 +1,605 @@
|
|||
"""Rubrik LiteLLM Plugin for tool blocking and batch logging."""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import random
|
||||
import time
|
||||
import urllib.parse
|
||||
import uuid
|
||||
from collections import Counter
|
||||
from typing import TYPE_CHECKING, Any, Literal, Optional
|
||||
|
||||
import httpx
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_batch_logger import CustomBatchLogger
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
ModifyResponseException,
|
||||
)
|
||||
from litellm.litellm_core_utils.core_helpers import safe_deep_copy
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import (
|
||||
ChatCompletionMessageToolCall,
|
||||
Function,
|
||||
GenericGuardrailAPIInputs,
|
||||
StandardLoggingPayload,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
Logging as LiteLLMLoggingObj,
|
||||
)
|
||||
|
||||
_ENDPOINT_ANTHROPIC_MESSAGES = "/v1/messages"
|
||||
_WEBHOOK_PATH_TOOL_BLOCKING = "/v1/after_completion/openai/v1"
|
||||
_WEBHOOK_PATH_LOGGING_BATCH = "/v1/litellm/batch"
|
||||
_MAX_QUEUE_SIZE = 10_000
|
||||
_DROP_WARNING_INTERVAL_SECONDS = 60.0
|
||||
|
||||
|
||||
class _MalformedToolBlockingResponseError(Exception):
|
||||
"""Raised when the tool blocking service returns a structurally invalid
|
||||
response (e.g. empty ``choices``).
|
||||
|
||||
Distinct from transient network/HTTP errors so callers can surface a
|
||||
louder, misconfiguration-style log instead of treating it as a routine
|
||||
fail-open.
|
||||
"""
|
||||
|
||||
|
||||
class RubrikLogger(CustomGuardrail, CustomBatchLogger):
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
self.flush_lock = asyncio.Lock()
|
||||
kwargs.setdefault("guardrail_name", "rubrik")
|
||||
# `initialize_guardrail` always passes these kwargs explicitly, with
|
||||
# value `None` when the user omits `mode` / `default_on` from the
|
||||
# guardrail config. Coerce None (omitted) to the desired default
|
||||
# while preserving any explicit value the caller did set --
|
||||
# in particular `default_on=False` if the user wants the guardrail
|
||||
# off by default.
|
||||
kwargs["event_hook"] = kwargs.get("event_hook") or GuardrailEventHooks.post_call
|
||||
if kwargs.get("default_on") is None:
|
||||
kwargs["default_on"] = True
|
||||
super().__init__(
|
||||
flush_lock=self.flush_lock,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
verbose_logger.debug("initializing rubrik logger")
|
||||
|
||||
self.sampling_rate = 1.0
|
||||
rbrk_sampling_rate = os.getenv("RUBRIK_SAMPLING_RATE")
|
||||
if rbrk_sampling_rate is not None:
|
||||
try:
|
||||
parsed_rate = float(rbrk_sampling_rate.strip())
|
||||
self.sampling_rate = max(0.0, min(1.0, parsed_rate))
|
||||
if parsed_rate != self.sampling_rate:
|
||||
verbose_logger.warning(
|
||||
f"RUBRIK_SAMPLING_RATE={parsed_rate} clamped to "
|
||||
f"{self.sampling_rate}"
|
||||
)
|
||||
except ValueError:
|
||||
verbose_logger.warning(
|
||||
f"Invalid RUBRIK_SAMPLING_RATE: {rbrk_sampling_rate!r}, using 1.0"
|
||||
)
|
||||
|
||||
self.key = api_key or os.getenv("RUBRIK_API_KEY")
|
||||
if not self.key:
|
||||
verbose_logger.warning(
|
||||
"Rubrik: No API key configured. Requests will be unauthenticated."
|
||||
)
|
||||
_batch_size = os.getenv("RUBRIK_BATCH_SIZE")
|
||||
|
||||
if _batch_size:
|
||||
try:
|
||||
self.batch_size = int(_batch_size)
|
||||
except ValueError:
|
||||
verbose_logger.warning(
|
||||
f"Invalid RUBRIK_BATCH_SIZE: {_batch_size!r}, using default"
|
||||
)
|
||||
|
||||
# Cap the in-memory retry queue so a Rubrik webhook outage cannot let
|
||||
# authenticated traffic accumulate prompt/response payloads until the
|
||||
# proxy runs out of memory. Once the cap is reached, oldest events are
|
||||
# dropped to make room for fresh ones (drop-oldest backpressure).
|
||||
self.max_queue_size = _MAX_QUEUE_SIZE
|
||||
self._dropped_since_warning = 0
|
||||
self._last_drop_warning_time = 0.0
|
||||
|
||||
_webhook_url = api_base or os.getenv("RUBRIK_WEBHOOK_URL")
|
||||
|
||||
if _webhook_url is None:
|
||||
raise ValueError(
|
||||
"Rubrik webhook URL not configured. "
|
||||
"Set RUBRIK_WEBHOOK_URL or pass api_base."
|
||||
)
|
||||
|
||||
_webhook_url = _webhook_url.rstrip("/").removesuffix("/v1")
|
||||
self.tool_blocking_endpoint = f"{_webhook_url}{_WEBHOOK_PATH_TOOL_BLOCKING}"
|
||||
self.logging_endpoint = f"{_webhook_url}{_WEBHOOK_PATH_LOGGING_BATCH}"
|
||||
|
||||
self.async_httpx_client = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.LoggingCallback
|
||||
)
|
||||
|
||||
self.tool_blocking_client = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.LoggingCallback,
|
||||
params={"timeout": httpx.Timeout(5.0, connect=2.0)},
|
||||
)
|
||||
|
||||
self._headers: dict[str, str] = {"Content-Type": "application/json"}
|
||||
if self.key:
|
||||
self._headers["Authorization"] = f"Bearer {self.key}"
|
||||
|
||||
# Periodic flush is started lazily on the first log event so that
|
||||
# low-traffic deployments still get their batches drained even when the
|
||||
# logger is instantiated outside a running event loop (sync init).
|
||||
self._flush_task: Optional[asyncio.Task[Any]] = (
|
||||
self._start_periodic_flush_task()
|
||||
)
|
||||
|
||||
def _start_periodic_flush_task(self) -> Optional[asyncio.Task[Any]]:
|
||||
"""Start the periodic flush task only when an event loop is already running."""
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
verbose_logger.debug(
|
||||
"Rubrik logger init: no running event loop, "
|
||||
"periodic flush will start on first log event."
|
||||
)
|
||||
return None
|
||||
return loop.create_task(self.periodic_flush())
|
||||
|
||||
def _ensure_periodic_flush_task(self) -> None:
|
||||
# Synchronous helper: in asyncio's cooperative model there is no await
|
||||
# between the check and assignment, so two callers cannot race here.
|
||||
if self._flush_task is None or self._flush_task.done():
|
||||
self._flush_task = self._start_periodic_flush_task()
|
||||
|
||||
async def aclose(self):
|
||||
"""Close the dedicated HTTP clients used by this logger."""
|
||||
# Cancel the periodic flush task before closing the HTTP clients so
|
||||
# the loop doesn't wake up and try to POST via a closed client.
|
||||
if self._flush_task is not None and not self._flush_task.done():
|
||||
self._flush_task.cancel()
|
||||
try:
|
||||
await self._flush_task
|
||||
except (asyncio.CancelledError, Exception):
|
||||
pass
|
||||
self._flush_task = None
|
||||
await self.tool_blocking_client.close()
|
||||
await self.async_httpx_client.close()
|
||||
|
||||
# -- Guardrail hook --------------------------------------------------------
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
"""Validate tool calls against the blocking service (fail-open)."""
|
||||
if input_type != "response":
|
||||
return inputs
|
||||
|
||||
tool_calls = inputs.get("tool_calls")
|
||||
if not tool_calls:
|
||||
return inputs
|
||||
|
||||
try:
|
||||
return await self._check_tool_calls(
|
||||
inputs, tool_calls, request_data, logging_obj
|
||||
)
|
||||
except ModifyResponseException:
|
||||
raise
|
||||
except _MalformedToolBlockingResponseError as e:
|
||||
# Distinct from transient errors: the service responded but the
|
||||
# payload was structurally invalid, which usually indicates a
|
||||
# misconfigured webhook or a breaking change in its response
|
||||
# format. Log loudly so operators notice their tool-blocking
|
||||
# policy is not actually being enforced.
|
||||
verbose_logger.critical(
|
||||
"Tool blocking service returned a malformed response: %s. "
|
||||
"Tool calls are NOT being checked -- verify the webhook "
|
||||
"configuration. Returning original response unchanged.",
|
||||
e,
|
||||
exc_info=True,
|
||||
)
|
||||
return inputs
|
||||
except Exception as e:
|
||||
verbose_logger.error(
|
||||
f"Tool blocking hook failed: {e}. "
|
||||
"Returning original response unchanged.",
|
||||
exc_info=True,
|
||||
)
|
||||
return inputs
|
||||
|
||||
async def _check_tool_calls(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
tool_calls: Any,
|
||||
request_data: dict,
|
||||
logging_obj: Optional["LiteLLMLoggingObj"],
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
"""Send tool calls to blocking service, raise if any are blocked."""
|
||||
message_tool_calls = self._normalize_tool_calls(tool_calls)
|
||||
|
||||
call_details = (
|
||||
getattr(logging_obj, "model_call_details", {}) if logging_obj else {}
|
||||
)
|
||||
response = request_data.get("response")
|
||||
request_id = getattr(response, "id", None) if response else None
|
||||
if logging_obj and not call_details:
|
||||
verbose_logger.warning(
|
||||
"Rubrik: logging_obj present but model_call_details is empty "
|
||||
"-- request context will be missing"
|
||||
)
|
||||
|
||||
response_data = self._build_tool_call_payload(message_tool_calls, request_id)
|
||||
req_data = self._extract_request_data(call_details)
|
||||
|
||||
service_response = await self._post_to_tool_blocking_service(
|
||||
response_data, req_data
|
||||
)
|
||||
blocked_explanation = self._extract_blocked_tools(
|
||||
service_response, message_tool_calls
|
||||
)
|
||||
|
||||
if blocked_explanation is not None:
|
||||
model = self._resolve_model(request_data, call_details)
|
||||
raise ModifyResponseException(
|
||||
message=blocked_explanation,
|
||||
model=model,
|
||||
request_data=request_data,
|
||||
guardrail_name=self.guardrail_name,
|
||||
)
|
||||
|
||||
return inputs
|
||||
|
||||
@staticmethod
|
||||
def _normalize_tool_calls(tool_calls: Any) -> list[ChatCompletionMessageToolCall]:
|
||||
"""Convert tool_calls from inputs to ChatCompletionMessageToolCall objects."""
|
||||
result = []
|
||||
for tc in tool_calls:
|
||||
if isinstance(tc, ChatCompletionMessageToolCall):
|
||||
result.append(tc)
|
||||
elif isinstance(tc, dict):
|
||||
func = tc.get("function", {})
|
||||
result.append(
|
||||
ChatCompletionMessageToolCall(
|
||||
id=tc.get("id", ""),
|
||||
type=tc.get("type", "function"),
|
||||
function=Function(
|
||||
name=func.get("name", ""),
|
||||
arguments=func.get("arguments", ""),
|
||||
),
|
||||
)
|
||||
)
|
||||
elif hasattr(tc, "id") and hasattr(tc, "function"):
|
||||
result.append(
|
||||
ChatCompletionMessageToolCall(
|
||||
id=tc.id or "",
|
||||
type=getattr(tc, "type", None) or "function",
|
||||
function=tc.function,
|
||||
)
|
||||
)
|
||||
else:
|
||||
raise TypeError(
|
||||
f"Cannot normalize tool_call of type {type(tc).__name__}"
|
||||
)
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _build_tool_call_payload(
|
||||
tool_calls: list[ChatCompletionMessageToolCall],
|
||||
request_id: str | None,
|
||||
) -> dict[str, Any]:
|
||||
"""Build a full OpenAI ChatCompletion-format dict for the blocking service."""
|
||||
return {
|
||||
"id": request_id or f"chatcmpl-{uuid.uuid4()}",
|
||||
"object": "chat.completion",
|
||||
"created": int(time.time()),
|
||||
"model": "",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
tc.model_dump(exclude_none=True) for tc in tool_calls
|
||||
],
|
||||
},
|
||||
"finish_reason": "tool_calls",
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _extract_request_data(call_details: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Extract original request data from model_call_details."""
|
||||
if not call_details:
|
||||
return {}
|
||||
litellm_params = call_details.get("litellm_params", {}) or {}
|
||||
return {
|
||||
"messages": call_details.get("messages"),
|
||||
"model": call_details.get("model"),
|
||||
"proxy_server_request": RubrikLogger._sanitize_proxy_server_request(
|
||||
litellm_params.get("proxy_server_request")
|
||||
),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _sanitize_proxy_server_request(proxy_server_request: Any) -> Any:
|
||||
"""Allowlist only routing fields (``url``, ``method``) when forwarding
|
||||
``proxy_server_request`` to the external Rubrik webhook, dropping
|
||||
inbound ``headers`` (Authorization, Cookie, x-api-key, ...) and the raw
|
||||
request ``body`` so proxy credentials are not exfiltrated."""
|
||||
if not isinstance(proxy_server_request, dict):
|
||||
return proxy_server_request
|
||||
return {
|
||||
key: proxy_server_request[key]
|
||||
for key in ("url", "method")
|
||||
if key in proxy_server_request
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _resolve_model(
|
||||
request_data: dict[str, Any], call_details: dict[str, Any]
|
||||
) -> str:
|
||||
"""Get the model name for the ModifyResponseException."""
|
||||
response = request_data.get("response")
|
||||
if response and hasattr(response, "model"):
|
||||
return response.model or "unknown"
|
||||
return call_details.get("model", "unknown")
|
||||
|
||||
# -- Logging hooks ---------------------------------------------------------
|
||||
|
||||
async def _prepare_log_payload(
|
||||
self, kwargs: dict, event_type: str
|
||||
) -> StandardLoggingPayload | None:
|
||||
"""Shared logic for success and failure logging."""
|
||||
if random.random() > self.sampling_rate:
|
||||
verbose_logger.debug(
|
||||
f"Skipping Rubrik {event_type} logging "
|
||||
f"(sampling_rate={self.sampling_rate})"
|
||||
)
|
||||
return None
|
||||
|
||||
# Deep-copy so mutations don't affect other callbacks sharing this object
|
||||
standard_logging_payload: StandardLoggingPayload = safe_deep_copy(
|
||||
kwargs["standard_logging_object"]
|
||||
)
|
||||
|
||||
# For Anthropic /v1/messages requests, LiteLLM creates a separate
|
||||
# ModelResponse (with a generated chatcmpl-* id) for logging, which
|
||||
# differs from the original Anthropic msg-* id on the response dict.
|
||||
# Normalize to litellm_call_id so that the logging and tool-blocking
|
||||
# endpoints see the same request identifier.
|
||||
litellm_params = kwargs.get("litellm_params", {}) or {}
|
||||
proxy_request = litellm_params.get("proxy_server_request", {}) or {}
|
||||
url_path = urllib.parse.urlparse(proxy_request.get("url", "")).path
|
||||
if url_path.endswith(_ENDPOINT_ANTHROPIC_MESSAGES):
|
||||
_litellm_call_id = kwargs.get("litellm_call_id")
|
||||
if _litellm_call_id:
|
||||
standard_logging_payload["id"] = _litellm_call_id # type: ignore[literal-required]
|
||||
|
||||
if "system" in kwargs:
|
||||
system_prompt_msg_list = kwargs["system"]
|
||||
try:
|
||||
if system_prompt_msg_list:
|
||||
system_scaffold = {
|
||||
"role": "system",
|
||||
"content": system_prompt_msg_list,
|
||||
}
|
||||
if isinstance(standard_logging_payload["messages"], list):
|
||||
standard_logging_payload["messages"].insert(0, system_scaffold)
|
||||
elif isinstance(standard_logging_payload["messages"], (dict, str)):
|
||||
standard_logging_payload["messages"] = [
|
||||
system_scaffold,
|
||||
standard_logging_payload["messages"],
|
||||
]
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
f"Rubrik: failed to prepend system prompt: {e}",
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
return standard_logging_payload
|
||||
|
||||
async def _enqueue_log_event(self, kwargs: dict, event_type: str):
|
||||
try:
|
||||
self._ensure_periodic_flush_task()
|
||||
payload = await self._prepare_log_payload(kwargs, event_type)
|
||||
if payload is None:
|
||||
return
|
||||
|
||||
self.log_queue.append(payload)
|
||||
self._enforce_max_queue_size()
|
||||
|
||||
if len(self.log_queue) >= self.batch_size:
|
||||
await self.flush_queue()
|
||||
except Exception as e:
|
||||
verbose_logger.error(
|
||||
f"Rubrik {event_type} logging hook failed: {e}. "
|
||||
"Skipping logging for this event.",
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
def _enforce_max_queue_size(self) -> None:
|
||||
overflow = len(self.log_queue) - self.max_queue_size
|
||||
if overflow <= 0:
|
||||
return
|
||||
del self.log_queue[:overflow]
|
||||
self._dropped_since_warning += overflow
|
||||
now = time.time()
|
||||
if now - self._last_drop_warning_time >= _DROP_WARNING_INTERVAL_SECONDS:
|
||||
verbose_logger.warning(
|
||||
"Rubrik: log queue exceeded max_queue_size=%s; dropped %s "
|
||||
"oldest events since the last warning. The Rubrik webhook may "
|
||||
"be unhealthy or undersized for current traffic.",
|
||||
self.max_queue_size,
|
||||
self._dropped_since_warning,
|
||||
)
|
||||
self._dropped_since_warning = 0
|
||||
self._last_drop_warning_time = now
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
await self._enqueue_log_event(kwargs, "success")
|
||||
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
await self._enqueue_log_event(kwargs, "failure")
|
||||
|
||||
# -- Batch logging ---------------------------------------------------------
|
||||
|
||||
async def _log_batch_to_rubrik(self, data):
|
||||
# NOTE: this method intentionally re-raises on failure so the parent
|
||||
# CustomBatchLogger.flush_queue keeps the unsent events in the queue
|
||||
# for the next flush attempt instead of silently dropping them.
|
||||
try:
|
||||
response = await self.async_httpx_client.post(
|
||||
url=self.logging_endpoint,
|
||||
json=data,
|
||||
headers=self._headers,
|
||||
)
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
verbose_logger.exception(
|
||||
f"Rubrik HTTP Error: {e.response.status_code} - {e.response.text}"
|
||||
)
|
||||
raise
|
||||
except Exception:
|
||||
verbose_logger.exception("Rubrik Layer Error")
|
||||
raise
|
||||
|
||||
async def async_send_batch(self):
|
||||
"""Handles sending batches of responses to Rubrik.
|
||||
|
||||
Note: the canonical flush path is :meth:`flush_queue`, which takes a
|
||||
single snapshot used for both sending and queue draining. This method
|
||||
is kept for direct callers / tests; it intentionally does NOT remove
|
||||
events from the queue.
|
||||
"""
|
||||
if not self.log_queue:
|
||||
return
|
||||
|
||||
log_queue_snapshot = list(self.log_queue)
|
||||
verbose_logger.debug(
|
||||
"Rubrik: Flushing batch of %s events", len(log_queue_snapshot)
|
||||
)
|
||||
await self._log_batch_to_rubrik(
|
||||
data=log_queue_snapshot,
|
||||
)
|
||||
|
||||
async def flush_queue(self):
|
||||
"""Snapshot, send, and drain in one consistent step.
|
||||
|
||||
Overrides the base implementation so the same snapshot drives both
|
||||
the HTTP send and the queue truncation. This avoids the subtle
|
||||
coupling where the base class captures `len(self.log_queue)`
|
||||
separately from the snapshot taken inside `async_send_batch`,
|
||||
which could otherwise drift in a future refactor and cause
|
||||
duplicate deliveries to Rubrik.
|
||||
"""
|
||||
if self.flush_lock is None:
|
||||
return
|
||||
|
||||
async with self.flush_lock:
|
||||
if not self.log_queue:
|
||||
return
|
||||
snapshot = list(self.log_queue)
|
||||
verbose_logger.debug("Rubrik: Flushing batch of %s events", len(snapshot))
|
||||
try:
|
||||
await self._log_batch_to_rubrik(data=snapshot)
|
||||
except Exception:
|
||||
# Already logged with traceback inside _log_batch_to_rubrik.
|
||||
# Preserve the in-flight events for retry on the next flush.
|
||||
return
|
||||
del self.log_queue[: len(snapshot)]
|
||||
self.last_flush_time = time.time()
|
||||
|
||||
# -- Tool blocking service -------------------------------------------------
|
||||
|
||||
async def _post_to_tool_blocking_service(
|
||||
self,
|
||||
response_data: dict[str, Any],
|
||||
request_data: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""Post a payload to the tool blocking service and return the response.
|
||||
|
||||
Args:
|
||||
response_data: The OpenAI-formatted response payload to send.
|
||||
request_data: Original LLM request data to include alongside
|
||||
the response for additional context. Empty dict if unavailable.
|
||||
|
||||
Raises:
|
||||
Exception: If the service is unavailable or returns an error.
|
||||
"""
|
||||
envelope = {
|
||||
"request": request_data,
|
||||
"response": response_data,
|
||||
}
|
||||
verbose_logger.debug(
|
||||
f"Sending request to tool blocking service: "
|
||||
f"{self.tool_blocking_endpoint}"
|
||||
)
|
||||
http_response = await self.tool_blocking_client.post(
|
||||
self.tool_blocking_endpoint,
|
||||
json=envelope,
|
||||
headers=self._headers,
|
||||
)
|
||||
http_response.raise_for_status()
|
||||
result: dict[str, Any] = http_response.json()
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _extract_blocked_tools(
|
||||
service_response: dict[str, Any],
|
||||
all_tool_calls: list[ChatCompletionMessageToolCall],
|
||||
) -> Optional[str]:
|
||||
"""Return the blocking explanation if any tool calls were blocked.
|
||||
|
||||
Compares the service response (which contains only allowed tools) against
|
||||
the full set of tool calls. Returns ``None`` if all tools are allowed, or
|
||||
the explanation string (prefixed with newlines) otherwise.
|
||||
|
||||
Expects service_response in OpenAI chat completion format:
|
||||
{"choices": [{"message": {"tool_calls": [...], "content": "..."}}]}
|
||||
"""
|
||||
choices = service_response.get("choices", [])
|
||||
if not choices:
|
||||
raise _MalformedToolBlockingResponseError(
|
||||
"Tool blocking service returned empty response"
|
||||
)
|
||||
|
||||
message = choices[0].get("message", {})
|
||||
returned_tool_calls = message.get("tool_calls") or []
|
||||
blocking_explanation = message.get("content", "")
|
||||
|
||||
allowed_id_counts: Counter = Counter(
|
||||
tc["id"]
|
||||
for tc in returned_tool_calls
|
||||
if isinstance(tc, dict) and tc.get("id")
|
||||
)
|
||||
required_id_counts: Counter = Counter(tc.id for tc in all_tool_calls if tc.id)
|
||||
|
||||
all_allowed = len(returned_tool_calls) >= len(all_tool_calls) and all(
|
||||
allowed_id_counts.get(tc_id, 0) >= count
|
||||
for tc_id, count in required_id_counts.items()
|
||||
)
|
||||
|
||||
if all_allowed:
|
||||
return None
|
||||
|
||||
explanation = blocking_explanation or "Tool call blocked by policy."
|
||||
return f"\n\n{explanation}"
|
||||
|
|
@ -1,8 +1,8 @@
|
|||
"""
|
||||
s3 Bucket Logging Integration
|
||||
|
||||
async_log_success_event: Processes the event, stores it in memory for DEFAULT_S3_FLUSH_INTERVAL_SECONDS seconds or until DEFAULT_S3_BATCH_SIZE and then flushes to s3
|
||||
async_log_failure_event: Processes the event, stores it in memory for DEFAULT_S3_FLUSH_INTERVAL_SECONDS seconds or until DEFAULT_S3_BATCH_SIZE and then flushes to s3
|
||||
async_log_success_event: Processes the event, stores it in memory for DEFAULT_S3_FLUSH_INTERVAL_SECONDS seconds or until DEFAULT_S3_BATCH_SIZE and then flushes to s3
|
||||
async_log_failure_event: Processes the event, stores it in memory for DEFAULT_S3_FLUSH_INTERVAL_SECONDS seconds or until DEFAULT_S3_BATCH_SIZE and then flushes to s3
|
||||
NOTE 1: S3 does not provide a BATCH PUT API endpoint, so we create tasks to upload each element individually
|
||||
"""
|
||||
|
||||
|
|
|
|||
|
|
@ -49,7 +49,6 @@ from litellm.types.interactions import InteractionEnvironment
|
|||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.utils import client
|
||||
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# Shared helpers #
|
||||
# ------------------------------------------------------------------ #
|
||||
|
|
|
|||
|
|
@ -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="...")
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ def get_supported_openai_params( # noqa: PLR0915
|
|||
request_type: Literal[
|
||||
"chat_completion", "embeddings", "transcription"
|
||||
] = "chat_completion",
|
||||
base_model: Optional[str] = None,
|
||||
) -> Optional[list]:
|
||||
"""
|
||||
Returns the supported openai params for a given model + provider
|
||||
|
|
@ -20,6 +21,11 @@ def get_supported_openai_params( # noqa: PLR0915
|
|||
get_supported_openai_params(model="anthropic.claude-3", custom_llm_provider="bedrock")
|
||||
```
|
||||
|
||||
Args:
|
||||
base_model: For Azure, the true underlying model (e.g. ``"azure/gpt-5.2"``)
|
||||
when the deployment name differs. Used for model-type detection so that
|
||||
non-standard deployment names route to the correct config.
|
||||
|
||||
Returns:
|
||||
- List if custom_llm_provider is mapped
|
||||
- None if unmapped
|
||||
|
|
@ -32,17 +38,21 @@ def get_supported_openai_params( # noqa: PLR0915
|
|||
|
||||
if custom_llm_provider in LlmProvidersSet:
|
||||
provider_config = litellm.ProviderConfigManager.get_provider_chat_config(
|
||||
model=model, provider=LlmProviders(custom_llm_provider)
|
||||
model=model,
|
||||
provider=LlmProviders(custom_llm_provider),
|
||||
base_model=base_model,
|
||||
)
|
||||
elif custom_llm_provider.split("/")[0] in LlmProvidersSet:
|
||||
provider_config = litellm.ProviderConfigManager.get_provider_chat_config(
|
||||
model=model, provider=LlmProviders(custom_llm_provider.split("/")[0])
|
||||
model=model,
|
||||
provider=LlmProviders(custom_llm_provider.split("/")[0]),
|
||||
base_model=base_model,
|
||||
)
|
||||
else:
|
||||
provider_config = None
|
||||
|
||||
if provider_config and request_type == "chat_completion":
|
||||
return provider_config.get_supported_openai_params(model=model)
|
||||
return provider_config.get_supported_openai_params(model=base_model or model)
|
||||
|
||||
if custom_llm_provider == "bedrock":
|
||||
return litellm.AmazonConverseConfig().get_supported_openai_params(model=model)
|
||||
|
|
@ -130,16 +140,23 @@ def get_supported_openai_params( # noqa: PLR0915
|
|||
model=model
|
||||
)
|
||||
elif custom_llm_provider == "azure":
|
||||
if litellm.AzureOpenAIO1Config().is_o_series_model(model=model):
|
||||
_azure_detection_model = base_model or model
|
||||
if litellm.AzureOpenAIO1Config().is_o_series_model(
|
||||
model=_azure_detection_model
|
||||
):
|
||||
return litellm.AzureOpenAIO1Config().get_supported_openai_params(
|
||||
model=model
|
||||
model=_azure_detection_model
|
||||
)
|
||||
elif litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(model=model):
|
||||
elif litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(
|
||||
model=_azure_detection_model
|
||||
):
|
||||
return litellm.AzureOpenAIGPT5Config().get_supported_openai_params(
|
||||
model=model
|
||||
model=_azure_detection_model
|
||||
)
|
||||
else:
|
||||
return litellm.AzureOpenAIConfig().get_supported_openai_params(model=model)
|
||||
return litellm.AzureOpenAIConfig().get_supported_openai_params(
|
||||
model=_azure_detection_model
|
||||
)
|
||||
elif custom_llm_provider == "openrouter":
|
||||
return litellm.OpenrouterConfig().get_supported_openai_params(model=model)
|
||||
elif custom_llm_provider == "vercel_ai_gateway":
|
||||
|
|
|
|||
|
|
@ -994,10 +994,8 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
try:
|
||||
# [Non-blocking Extra Debug Information in metadata]
|
||||
if turn_off_message_logging is True:
|
||||
_metadata["raw_request"] = (
|
||||
"redacted by litellm. \
|
||||
_metadata["raw_request"] = "redacted by litellm. \
|
||||
'litellm.turn_off_message_logging=True'"
|
||||
)
|
||||
else:
|
||||
curl_command = self._get_request_curl_command(
|
||||
api_base=additional_args.get("api_base", ""),
|
||||
|
|
@ -1031,12 +1029,8 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
error=str(e),
|
||||
)
|
||||
)
|
||||
_metadata["raw_request"] = (
|
||||
"Unable to Log \
|
||||
raw request: {}".format(
|
||||
str(e)
|
||||
)
|
||||
)
|
||||
_metadata["raw_request"] = "Unable to Log \
|
||||
raw request: {}".format(str(e))
|
||||
if getattr(self, "logger_fn", None) and callable(self.logger_fn):
|
||||
try:
|
||||
self.logger_fn(
|
||||
|
|
@ -1769,9 +1763,12 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
self.model_call_details["response_cost"] = 0.0
|
||||
elif "response_cost" in hidden_params:
|
||||
self.model_call_details["response_cost"] = hidden_params["response_cost"]
|
||||
elif self.model_call_details.get("response_cost") is not None:
|
||||
elif (
|
||||
existing_cost := self.model_call_details.get("response_cost")
|
||||
) is not None and existing_cost != 0:
|
||||
# Preserve response_cost if already calculated (e.g., by pass-through
|
||||
# handlers like Gemini/Vertex which call completion_cost directly)
|
||||
# handlers like Gemini/Vertex which call completion_cost directly).
|
||||
# Do not preserve 0 from failure_handler on intermediate router retries.
|
||||
pass
|
||||
else:
|
||||
self.model_call_details["response_cost"] = self._response_cost_calculator(
|
||||
|
|
@ -5143,13 +5140,17 @@ class StandardLoggingPayloadSetup:
|
|||
) -> StandardLoggingPayloadErrorInformation:
|
||||
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
|
||||
|
||||
# Check for 'code' first (used by ProxyException), then fall back to 'status_code' (used by LiteLLM exceptions)
|
||||
# Ensure error_code is always a string for Prisma Python JSON field compatibility
|
||||
# ProxyException uses .code, LiteLLM exceptions use .status_code,
|
||||
# httpx.HTTPStatusError exposes status only as .response.status_code.
|
||||
# Stringified for Prisma JSON compatibility.
|
||||
error_code_attr = getattr(original_exception, "code", None)
|
||||
if error_code_attr is not None and str(error_code_attr) not in ("", "None"):
|
||||
error_status: str = str(error_code_attr)
|
||||
else:
|
||||
status_code_attr = getattr(original_exception, "status_code", None)
|
||||
if status_code_attr is None:
|
||||
response_attr = getattr(original_exception, "response", None)
|
||||
status_code_attr = getattr(response_attr, "status_code", None)
|
||||
error_status = str(status_code_attr) if status_code_attr is not None else ""
|
||||
error_class: str = (
|
||||
str(original_exception.__class__.__name__) if original_exception else ""
|
||||
|
|
|
|||
|
|
@ -1204,12 +1204,8 @@ def get_last_user_message(messages: List[AllMessageValues]) -> Optional[str]:
|
|||
{"role": "assistant", "content": "I'm good, thank you!"},
|
||||
{"role": "user", "content": "What is the weather in Tokyo?"},
|
||||
]
|
||||
get_user_prompt(messages) -> "What is the weather in Tokyo?"
|
||||
get_last_user_message(messages) -> "What is the weather in Tokyo?"
|
||||
"""
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
convert_content_list_to_str,
|
||||
)
|
||||
|
||||
if not messages:
|
||||
return None
|
||||
|
||||
|
|
|
|||
|
|
@ -5590,9 +5590,7 @@ def default_response_schema_prompt(response_schema: dict) -> str:
|
|||
prompt_str = """Use this JSON schema:
|
||||
```json
|
||||
{}
|
||||
```""".format(
|
||||
response_schema
|
||||
)
|
||||
```""".format(response_schema)
|
||||
return prompt_str
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1476,7 +1476,7 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
for choice in choices:
|
||||
if choice.delta.content is not None and len(choice.delta.content) > 0:
|
||||
text += choice.delta.content
|
||||
if choice.delta.tool_calls is not None:
|
||||
if choice.delta.tool_calls:
|
||||
partial_json = ""
|
||||
for tool in choice.delta.tool_calls:
|
||||
if (
|
||||
|
|
|
|||
|
|
@ -13,7 +13,6 @@ from typing import Any, AsyncIterator, Dict, List, Optional, cast
|
|||
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SSE parsing helpers (module-level to keep the class lean)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -293,6 +293,12 @@ async def anthropic_messages(
|
|||
api_base=api_base,
|
||||
client=client,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
# messages were already empty-text-block sanitized at the top of this
|
||||
# function and are NOT reassigned before this dispatch, so the handler
|
||||
# can skip its (otherwise redundant) second full-messages scan. Passed
|
||||
# explicitly (not via **kwargs) so it only affects this direct
|
||||
# dispatch -- interceptor / sync entry points still sanitize.
|
||||
_litellm_messages_presanitized=True,
|
||||
**kwargs,
|
||||
)
|
||||
ctx = contextvars.copy_context()
|
||||
|
|
@ -351,10 +357,14 @@ def anthropic_messages_handler(
|
|||
"""
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
# Sanitize empty text blocks here too so the sync entry point
|
||||
# Sanitize empty text blocks so the sync entry point
|
||||
# (litellm.messages.create -> anthropic_messages_handler) gets the same
|
||||
# protection as the async wrapper. Idempotent when called twice.
|
||||
messages = strip_empty_text_blocks_from_anthropic_messages(messages)
|
||||
# protection as the async wrapper. The async wrapper already sanitized and
|
||||
# does not reassign messages before dispatch, so it sets
|
||||
# ``_litellm_messages_presanitized`` to skip this redundant second
|
||||
# full-messages scan. Pop it so it never leaks into provider params.
|
||||
if not kwargs.pop("_litellm_messages_presanitized", False):
|
||||
messages = strip_empty_text_blocks_from_anthropic_messages(messages)
|
||||
|
||||
metadata = validate_anthropic_api_metadata(metadata)
|
||||
|
||||
|
|
|
|||
|
|
@ -312,7 +312,10 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
|
|||
)
|
||||
|
||||
####### get required params for all anthropic messages requests ######
|
||||
verbose_logger.debug(f"TRANSFORMATION DEBUG - Messages: {messages}")
|
||||
# Lazy %s: the f-string previously stringified the entire messages
|
||||
# payload on every request regardless of log level (a full scan of the
|
||||
# request body on the hot path). Defer it to when DEBUG is enabled.
|
||||
verbose_logger.debug("TRANSFORMATION DEBUG - Messages: %s", messages)
|
||||
|
||||
# Auto-strip advisor blocks from history if advisor tool is absent.
|
||||
# Prevents Anthropic 400: advisor_tool_result in history requires advisor tool.
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
from typing import Any, Dict, List, cast, get_type_hints
|
||||
from functools import lru_cache
|
||||
from typing import Any, Dict, FrozenSet, List, cast, get_type_hints
|
||||
|
||||
from litellm.types.llms.anthropic import AnthropicMessagesRequestOptionalParams
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
|
|
@ -6,6 +7,18 @@ from litellm.types.llms.anthropic_messages.anthropic_response import (
|
|||
)
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _anthropic_messages_optional_param_keys() -> FrozenSet[str]:
|
||||
"""
|
||||
Valid AnthropicMessagesRequestOptionalParams keys.
|
||||
|
||||
``typing.get_type_hints`` is ~80us/call and this TypedDict is static, so
|
||||
resolving it once per process instead of once per request removes a fixed
|
||||
full-pass cost from the /v1/messages request-parse path.
|
||||
"""
|
||||
return frozenset(get_type_hints(AnthropicMessagesRequestOptionalParams).keys())
|
||||
|
||||
|
||||
class AnthropicMessagesRequestUtils:
|
||||
@staticmethod
|
||||
def get_requested_anthropic_messages_optional_param(
|
||||
|
|
@ -20,7 +33,7 @@ class AnthropicMessagesRequestUtils:
|
|||
Returns:
|
||||
AnthropicMessagesRequestOptionalParams instance with only the valid parameters
|
||||
"""
|
||||
valid_keys = get_type_hints(AnthropicMessagesRequestOptionalParams).keys()
|
||||
valid_keys = _anthropic_messages_optional_param_keys()
|
||||
filtered_params = {
|
||||
k: v for k, v in params.items() if k in valid_keys and v is not None
|
||||
}
|
||||
|
|
|
|||
3
litellm/llms/azure/audio_transcription/__init__.py
Normal file
3
litellm/llms/azure/audio_transcription/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from .transformation import AzureSpeechAudioTranscriptionConfig
|
||||
|
||||
__all__ = ["AzureSpeechAudioTranscriptionConfig"]
|
||||
224
litellm/llms/azure/audio_transcription/transformation.py
Normal file
224
litellm/llms/azure/audio_transcription/transformation.py
Normal file
|
|
@ -0,0 +1,224 @@
|
|||
"""
|
||||
Azure AI Speech (Cognitive Services) speech-to-text transformation.
|
||||
|
||||
Maps OpenAI-compatible audio transcription calls to Azure Speech REST
|
||||
recognition for short audio.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
from urllib.parse import urlencode, urlparse
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.litellm_core_utils.audio_utils.utils import process_audio_file
|
||||
from litellm.llms.base_llm.audio_transcription.transformation import (
|
||||
AudioTranscriptionRequestData,
|
||||
BaseAudioTranscriptionConfig,
|
||||
)
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
OpenAIAudioTranscriptionOptionalParams,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.utils import FileTypes, TranscriptionResponse
|
||||
|
||||
|
||||
class AzureSpeechAudioTranscriptionException(BaseLLMException):
|
||||
pass
|
||||
|
||||
|
||||
class AzureSpeechAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
|
||||
"""
|
||||
Configuration for Azure AI Speech (Cognitive Services) STT.
|
||||
|
||||
Reference:
|
||||
https://learn.microsoft.com/en-us/azure/ai-services/speech-service/rest-speech-to-text-short
|
||||
"""
|
||||
|
||||
COGNITIVE_SERVICES_DOMAIN = "api.cognitive.microsoft.com"
|
||||
STT_SPEECH_DOMAIN = "stt.speech.microsoft.com"
|
||||
STT_ENDPOINT_PATH = "/speech/recognition/conversation/cognitiveservices/v1"
|
||||
DEFAULT_LANGUAGE = "en-US"
|
||||
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> List[OpenAIAudioTranscriptionOptionalParams]:
|
||||
return ["language", "response_format"]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
supported_params = self.get_supported_openai_params(model=model)
|
||||
for key, value in non_default_params.items():
|
||||
if key in supported_params:
|
||||
optional_params[key] = value
|
||||
return optional_params
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
api_key = api_key or get_secret_str("AZURE_SPEECH_API_KEY")
|
||||
if not api_key:
|
||||
raise AzureSpeechAudioTranscriptionException(
|
||||
message="api_key is required for Azure AI Speech transcription.",
|
||||
status_code=401,
|
||||
)
|
||||
|
||||
validated_headers = headers.copy()
|
||||
validated_headers["Ocp-Apim-Subscription-Key"] = api_key
|
||||
validated_headers["Content-Type"] = validated_headers.get(
|
||||
"Content-Type", "audio/wav"
|
||||
)
|
||||
validated_headers["Accept"] = "application/json"
|
||||
return validated_headers
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
api_base = api_base or get_secret_str("AZURE_SPEECH_API_BASE")
|
||||
if api_base is None:
|
||||
raise AzureSpeechAudioTranscriptionException(
|
||||
message=(
|
||||
"api_base is required for Azure AI Speech transcription. "
|
||||
"Use a Cognitive Services endpoint like "
|
||||
"https://{region}.api.cognitive.microsoft.com or an STT "
|
||||
"endpoint like https://{region}.stt.speech.microsoft.com."
|
||||
),
|
||||
status_code=400,
|
||||
)
|
||||
|
||||
base_url = self._resolve_stt_base_url(api_base=api_base)
|
||||
query_params = {
|
||||
"language": optional_params.get("language", self.DEFAULT_LANGUAGE),
|
||||
"format": self._get_azure_response_format(
|
||||
optional_params.get("response_format")
|
||||
),
|
||||
}
|
||||
return f"{base_url}{self.STT_ENDPOINT_PATH}?{urlencode(query_params)}"
|
||||
|
||||
def transform_audio_transcription_request(
|
||||
self,
|
||||
model: str,
|
||||
audio_file: FileTypes,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
) -> AudioTranscriptionRequestData:
|
||||
processed_audio = process_audio_file(audio_file)
|
||||
return AudioTranscriptionRequestData(
|
||||
data=processed_audio.file_content,
|
||||
files=None,
|
||||
content_type=processed_audio.content_type,
|
||||
)
|
||||
|
||||
def transform_audio_transcription_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
) -> TranscriptionResponse:
|
||||
response_json = raw_response.json()
|
||||
recognition_status = response_json.get("RecognitionStatus")
|
||||
if recognition_status is not None and recognition_status != "Success":
|
||||
raise AzureSpeechAudioTranscriptionException(
|
||||
message=(
|
||||
"Azure AI Speech transcription failed with "
|
||||
f"RecognitionStatus={recognition_status}."
|
||||
),
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
|
||||
text = self._extract_text(response_json)
|
||||
response = TranscriptionResponse(text=text)
|
||||
response._hidden_params = response_json
|
||||
return response
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
) -> BaseLLMException:
|
||||
return AzureSpeechAudioTranscriptionException(
|
||||
message=error_message,
|
||||
status_code=status_code,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
def _resolve_stt_base_url(self, api_base: str) -> str:
|
||||
api_base = api_base.rstrip("/")
|
||||
parsed_url = urlparse(api_base)
|
||||
hostname = parsed_url.hostname or ""
|
||||
|
||||
if self._is_cognitive_services_endpoint(hostname=hostname):
|
||||
region = self._extract_region_from_hostname(
|
||||
hostname=hostname, domain=self.COGNITIVE_SERVICES_DOMAIN
|
||||
)
|
||||
return self._build_stt_base_url(region=region)
|
||||
|
||||
if self._is_stt_endpoint(hostname=hostname):
|
||||
return f"{parsed_url.scheme}://{hostname}"
|
||||
|
||||
if self._is_azure_openai_endpoint(hostname=hostname):
|
||||
raise AzureSpeechAudioTranscriptionException(
|
||||
message=(
|
||||
"Azure AI Speech transcription requires a Cognitive Services "
|
||||
"or STT Speech endpoint, not an Azure OpenAI endpoint."
|
||||
),
|
||||
status_code=400,
|
||||
)
|
||||
|
||||
return api_base
|
||||
|
||||
def _is_cognitive_services_endpoint(self, hostname: str) -> bool:
|
||||
return hostname == self.COGNITIVE_SERVICES_DOMAIN or hostname.endswith(
|
||||
f".{self.COGNITIVE_SERVICES_DOMAIN}"
|
||||
)
|
||||
|
||||
def _is_stt_endpoint(self, hostname: str) -> bool:
|
||||
return hostname == self.STT_SPEECH_DOMAIN or hostname.endswith(
|
||||
f".{self.STT_SPEECH_DOMAIN}"
|
||||
)
|
||||
|
||||
def _is_azure_openai_endpoint(self, hostname: str) -> bool:
|
||||
return hostname.endswith(".openai.azure.com")
|
||||
|
||||
def _extract_region_from_hostname(self, hostname: str, domain: str) -> str:
|
||||
if hostname.endswith(f".{domain}"):
|
||||
return hostname[: -len(f".{domain}")]
|
||||
return ""
|
||||
|
||||
def _build_stt_base_url(self, region: str) -> str:
|
||||
if region:
|
||||
return f"https://{region}.{self.STT_SPEECH_DOMAIN}"
|
||||
return f"https://{self.STT_SPEECH_DOMAIN}"
|
||||
|
||||
def _get_azure_response_format(self, response_format: Optional[str]) -> str:
|
||||
if response_format == "verbose_json":
|
||||
return "detailed"
|
||||
return "simple"
|
||||
|
||||
def _extract_text(self, response_json: Dict[str, Any]) -> str:
|
||||
if isinstance(response_json.get("DisplayText"), str):
|
||||
return response_json["DisplayText"]
|
||||
|
||||
nbest = response_json.get("NBest")
|
||||
if isinstance(nbest, list) and nbest:
|
||||
best = nbest[0]
|
||||
if isinstance(best, dict):
|
||||
return best.get("Display") or best.get("Lexical") or ""
|
||||
|
||||
return ""
|
||||
|
|
@ -239,7 +239,9 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
)
|
||||
|
||||
data = {"model": None, "messages": messages, **optional_params}
|
||||
elif litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(model=model):
|
||||
elif litellm.AzureOpenAIGPT5Config.is_model_gpt_5_model(
|
||||
model=litellm_params.get("base_model") or model
|
||||
):
|
||||
data = litellm.AzureOpenAIGPT5Config().transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
|
|
|
|||
|
|
@ -4,10 +4,10 @@ Support for o1 and o3 model families
|
|||
https://platform.openai.com/docs/guides/reasoning
|
||||
|
||||
Translations handled by LiteLLM:
|
||||
- modalities: image => drop param (if user opts in to dropping param)
|
||||
- role: system ==> translate to role 'user'
|
||||
- streaming => faked by LiteLLM
|
||||
- Tools, response_format => drop param (if user opts in to dropping param)
|
||||
- modalities: image => drop param (if user opts in to dropping param)
|
||||
- role: system ==> translate to role 'user'
|
||||
- streaming => faked by LiteLLM
|
||||
- Tools, response_format => drop param (if user opts in to dropping param)
|
||||
- Logprobs => drop param (if user opts in to dropping param)
|
||||
- Temperature => drop param (if user opts in to dropping param)
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Transformation logic from OpenAI /v1/embeddings format to Azure AI Cohere's /v1/embed.
|
||||
Transformation logic from OpenAI /v1/embeddings format to Azure AI Cohere's /v1/embed.
|
||||
|
||||
Why separate file? Make it easy to see how transformation works
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Translate between Cohere's `/rerank` format and Azure AI's `/rerank` format.
|
||||
Translate between Cohere's `/rerank` format and Azure AI's `/rerank` format.
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
|
|
|
|||
|
|
@ -54,6 +54,7 @@ class OCRUsageInfo(LiteLLMPydanticObjectBase):
|
|||
"""Usage information from OCR response."""
|
||||
|
||||
pages_processed: Optional[int] = None
|
||||
credits: Optional[float] = None
|
||||
doc_size_bytes: Optional[int] = None
|
||||
|
||||
model_config = {"extra": "allow"}
|
||||
|
|
|
|||
|
|
@ -44,6 +44,12 @@ else:
|
|||
# (e.g. "us-east-1", "eu-west-2", "us-gov-west-1", "cn-north-1").
|
||||
_VALID_AWS_REGION_PATTERN = re.compile(r"\A[a-z0-9-]+\Z")
|
||||
|
||||
# Regional STS hostnames, e.g. sts.eu-west-1.amazonaws.com or
|
||||
# vpce-xxx.sts.eu-west-1.vpce.amazonaws.com
|
||||
_STS_REGION_FROM_ENDPOINT_PATTERN = re.compile(
|
||||
r"(?:^|\.)sts(?:-fips)?\.([a-z0-9-]+)\.(?:amazonaws\.com(?:\.cn)?|vpce\.amazonaws\.com)"
|
||||
)
|
||||
|
||||
|
||||
class Boto3CredentialsInfo(BaseModel):
|
||||
credentials: Credentials
|
||||
|
|
@ -450,6 +456,24 @@ class BaseAWSLLM:
|
|||
model_id = BaseAWSLLM.encode_model_id(model_id=model_id)
|
||||
else:
|
||||
model_id = model
|
||||
# Strip LiteLLM routing prefixes (e.g. "bedrock/", "invoke/",
|
||||
# "bedrock/invoke/", "bedrock/converse/") that are not part of the
|
||||
# actual Bedrock model ID. The converse path already does this; the
|
||||
# invoke path must do the same so that ARN models such as
|
||||
# bedrock/arn:aws:bedrock:…:inference-profile/global.anthropic.…
|
||||
# are not forwarded verbatim to the Bedrock API, which would produce
|
||||
# a malformed URL and cause botocore's EventStreamBuffer to receive
|
||||
# a JSON error body instead of a binary event-stream — surfaced as a
|
||||
# misleading ChecksumMismatch (0x223a7b22 == ':{"').
|
||||
# Use strip_bedrock_routing_prefix (no break) so compound prefixes
|
||||
# like "bedrock/invoke/arn:..." are fully stripped in one call.
|
||||
from litellm.llms.bedrock.common_utils import strip_bedrock_routing_prefix
|
||||
|
||||
model_id = strip_bedrock_routing_prefix(model_id)
|
||||
# URL-encode ARNs so colons and slashes are safe in the URL path.
|
||||
if model_id.startswith("arn:"):
|
||||
model_id = BaseAWSLLM.encode_model_id(model_id=model_id)
|
||||
return model_id
|
||||
|
||||
model_id = model_id.replace("invoke/", "", 1)
|
||||
if provider == "llama" and "llama/" in model_id:
|
||||
|
|
@ -633,6 +657,40 @@ class BaseAWSLLM:
|
|||
"Region names must contain only lowercase letters, digits, and hyphens."
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _parse_sts_region_from_endpoint(
|
||||
aws_sts_endpoint: Optional[str],
|
||||
) -> Optional[str]:
|
||||
"""Extract region from sts.{region}.amazonaws.com or vpce-x.sts.{region}.vpce.amazonaws.com."""
|
||||
if not aws_sts_endpoint:
|
||||
return None
|
||||
host = urllib.parse.urlparse(aws_sts_endpoint).hostname or ""
|
||||
match = _STS_REGION_FROM_ENDPOINT_PATTERN.search(host)
|
||||
return match.group(1) if match else None
|
||||
|
||||
@staticmethod
|
||||
def _resolve_sts_region(aws_sts_endpoint: Optional[str] = None) -> Optional[str]:
|
||||
"""STS signing region: parsed from aws_sts_endpoint else AWS_REGION / AWS_DEFAULT_REGION."""
|
||||
return (
|
||||
BaseAWSLLM._parse_sts_region_from_endpoint(aws_sts_endpoint)
|
||||
or os.getenv("AWS_REGION")
|
||||
or os.getenv("AWS_DEFAULT_REGION")
|
||||
)
|
||||
|
||||
def _build_sts_client_kwargs(
|
||||
self,
|
||||
aws_sts_endpoint: Optional[str] = None,
|
||||
ssl_verify: Optional[Union[bool, str]] = None,
|
||||
) -> dict:
|
||||
"""STS client kwargs with aligned endpoint_url and region_name (SigV4)."""
|
||||
kwargs: dict = {"verify": self._get_ssl_verify(ssl_verify)}
|
||||
if aws_sts_endpoint is not None:
|
||||
kwargs["endpoint_url"] = aws_sts_endpoint
|
||||
sts_region = self._resolve_sts_region(aws_sts_endpoint)
|
||||
if sts_region is not None:
|
||||
kwargs["region_name"] = sts_region
|
||||
return kwargs
|
||||
|
||||
def get_aws_region_name_for_non_llm_api_calls(
|
||||
self,
|
||||
aws_region_name: Optional[str] = None,
|
||||
|
|
@ -787,11 +845,6 @@ class BaseAWSLLM:
|
|||
f"IN Web Identity Token: {aws_web_identity_token} | Role Name: {aws_role_name} | Session Name: {aws_session_name}"
|
||||
)
|
||||
|
||||
if aws_sts_endpoint is None:
|
||||
sts_endpoint = f"https://sts.{aws_region_name}.amazonaws.com"
|
||||
else:
|
||||
sts_endpoint = aws_sts_endpoint
|
||||
|
||||
oidc_token = get_secret(aws_web_identity_token)
|
||||
|
||||
if oidc_token is None:
|
||||
|
|
@ -800,13 +853,13 @@ class BaseAWSLLM:
|
|||
status_code=401,
|
||||
)
|
||||
|
||||
sts_client_kwargs = self._build_sts_client_kwargs(
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
with tracer.trace("boto3.client(sts)"):
|
||||
sts_client = boto3.client(
|
||||
"sts",
|
||||
region_name=aws_region_name,
|
||||
endpoint_url=sts_endpoint,
|
||||
verify=self._get_ssl_verify(ssl_verify),
|
||||
)
|
||||
sts_client = boto3.client("sts", **sts_client_kwargs)
|
||||
|
||||
# https://docs.aws.amazon.com/STS/latest/APIReference/API_AssumeRoleWithWebIdentity.html
|
||||
# https://boto3.amazonaws.com/v1/documentation/api/latest/reference/services/sts/client/assume_role_with_web_identity.html
|
||||
|
|
@ -847,7 +900,6 @@ class BaseAWSLLM:
|
|||
irsa_role_arn: str,
|
||||
aws_role_name: str,
|
||||
aws_session_name: str,
|
||||
region: str,
|
||||
web_identity_token_file: str,
|
||||
aws_external_id: Optional[str] = None,
|
||||
aws_sts_endpoint: Optional[str] = None,
|
||||
|
|
@ -862,12 +914,10 @@ class BaseAWSLLM:
|
|||
with open(web_identity_token_file, "r") as f:
|
||||
web_identity_token = f.read().strip()
|
||||
|
||||
irsa_sts_kwargs: dict = {
|
||||
"region_name": region,
|
||||
"verify": self._get_ssl_verify(ssl_verify),
|
||||
}
|
||||
if aws_sts_endpoint is not None:
|
||||
irsa_sts_kwargs["endpoint_url"] = aws_sts_endpoint
|
||||
irsa_sts_kwargs = self._build_sts_client_kwargs(
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
# Create an STS client without credentials
|
||||
with tracer.trace("boto3.client(sts) for manual IRSA"):
|
||||
|
|
@ -924,7 +974,6 @@ class BaseAWSLLM:
|
|||
self,
|
||||
aws_role_name: str,
|
||||
aws_session_name: str,
|
||||
region: str,
|
||||
aws_external_id: Optional[str] = None,
|
||||
aws_sts_endpoint: Optional[str] = None,
|
||||
ssl_verify: Optional[Union[bool, str]] = None,
|
||||
|
|
@ -932,12 +981,10 @@ class BaseAWSLLM:
|
|||
"""Handle same-account role assumption for IRSA."""
|
||||
import boto3
|
||||
|
||||
irsa_sts_kwargs: dict = {
|
||||
"region_name": region,
|
||||
"verify": self._get_ssl_verify(ssl_verify),
|
||||
}
|
||||
if aws_sts_endpoint is not None:
|
||||
irsa_sts_kwargs["endpoint_url"] = aws_sts_endpoint
|
||||
irsa_sts_kwargs = self._build_sts_client_kwargs(
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
|
||||
verbose_logger.debug("Same account role assumption, using automatic IRSA")
|
||||
with tracer.trace("boto3.client(sts) with automatic IRSA"):
|
||||
|
|
@ -1010,12 +1057,6 @@ class BaseAWSLLM:
|
|||
web_identity_token_file = os.getenv("AWS_WEB_IDENTITY_TOKEN_FILE")
|
||||
irsa_role_arn = os.getenv("AWS_ROLE_ARN")
|
||||
|
||||
region = (
|
||||
aws_region_name
|
||||
or os.getenv("AWS_REGION")
|
||||
or os.getenv("AWS_DEFAULT_REGION")
|
||||
)
|
||||
|
||||
# If we have IRSA environment variables and no explicit credentials,
|
||||
# we need to use the web identity token flow
|
||||
if (
|
||||
|
|
@ -1031,16 +1072,12 @@ class BaseAWSLLM:
|
|||
)
|
||||
|
||||
try:
|
||||
# Use passed-in region when set, else env, else default (align with AssumeRole path)
|
||||
region = region or "us-east-1"
|
||||
|
||||
# Check if we need to do cross-account role assumption
|
||||
if aws_role_name != irsa_role_arn:
|
||||
sts_response = self._handle_irsa_cross_account(
|
||||
irsa_role_arn,
|
||||
aws_role_name,
|
||||
aws_session_name,
|
||||
region,
|
||||
web_identity_token_file,
|
||||
aws_external_id,
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
|
|
@ -1050,7 +1087,6 @@ class BaseAWSLLM:
|
|||
sts_response = self._handle_irsa_same_account(
|
||||
aws_role_name,
|
||||
aws_session_name,
|
||||
region,
|
||||
aws_external_id,
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
|
|
@ -1074,11 +1110,10 @@ class BaseAWSLLM:
|
|||
|
||||
# In EKS/IRSA environments, use ambient credentials (no explicit keys needed)
|
||||
# This allows the web identity token to work automatically
|
||||
sts_client_kwargs: dict = {"verify": self._get_ssl_verify(ssl_verify)}
|
||||
if region is not None:
|
||||
sts_client_kwargs["region_name"] = region
|
||||
if aws_sts_endpoint is not None:
|
||||
sts_client_kwargs["endpoint_url"] = aws_sts_endpoint
|
||||
sts_client_kwargs = self._build_sts_client_kwargs(
|
||||
aws_sts_endpoint=aws_sts_endpoint,
|
||||
ssl_verify=ssl_verify,
|
||||
)
|
||||
if aws_access_key_id is None and aws_secret_access_key is None:
|
||||
with tracer.trace("boto3.client(sts)"):
|
||||
sts_client = boto3.client("sts", **sts_client_kwargs)
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ from typing import TYPE_CHECKING, Any, List, Optional
|
|||
import httpx
|
||||
|
||||
from litellm.anthropic_beta_headers_manager import filter_and_transform_beta_headers
|
||||
from litellm.litellm_core_utils.litellm_logging import verbose_logger
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
convert_to_anthropic_image_obj,
|
||||
)
|
||||
|
|
@ -22,6 +23,7 @@ from litellm.llms.bedrock.common_utils import (
|
|||
from litellm.types.llms.anthropic import ANTHROPIC_TOOL_SEARCH_BETA_HEADER
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse
|
||||
from litellm.utils import _supports_factory
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
|
@ -169,6 +171,24 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
|
|||
anthropic_request.pop("model", None)
|
||||
anthropic_request.pop("stream", None)
|
||||
anthropic_request.pop("output_format", None)
|
||||
if not (
|
||||
_supports_factory(
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
key="supports_output_config",
|
||||
)
|
||||
or AnthropicConfig._model_supports_effort_param(model)
|
||||
):
|
||||
if anthropic_request.pop("output_config", None) is not None:
|
||||
verbose_logger.warning(
|
||||
"Bedrock Invoke: stripping unsupported `output_config` for "
|
||||
"model=%s — neither `supports_output_config` nor any "
|
||||
"`supports_*_reasoning_effort` flag is set in "
|
||||
"model_prices_and_context_window.json. Add the capability "
|
||||
"flag to the model JSON entry if this model accepts "
|
||||
"`output_config`.",
|
||||
model,
|
||||
)
|
||||
if "anthropic_version" not in anthropic_request:
|
||||
anthropic_request["anthropic_version"] = self.anthropic_version
|
||||
|
||||
|
|
|
|||
|
|
@ -4,7 +4,6 @@ import litellm
|
|||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
|
||||
CLAUDE_PLATFORM_SERVICE_NAME: Literal["aws-external-anthropic"] = (
|
||||
"aws-external-anthropic"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Transformation logic from OpenAI /v1/embeddings format to Bedrock Amazon Titan G1 /invoke format.
|
||||
Transformation logic from OpenAI /v1/embeddings format to Bedrock Amazon Titan G1 /invoke format.
|
||||
|
||||
Why separate file? Make it easy to see how transformation works
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Transformation logic from OpenAI /v1/embeddings format to Bedrock Cohere /invoke format.
|
||||
Transformation logic from OpenAI /v1/embeddings format to Bedrock Cohere /invoke format.
|
||||
|
||||
Why separate file? Make it easy to see how transformation works
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -45,6 +45,7 @@ from litellm.types.router import GenericLiteLLMParams
|
|||
from litellm.types.utils import GenericStreamingChunk
|
||||
from litellm.types.utils import GenericStreamingChunk as GChunk
|
||||
from litellm.types.utils import ModelResponseStream
|
||||
from litellm.utils import _supports_factory
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
|
@ -557,7 +558,29 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
anthropic_messages_request=anthropic_messages_request,
|
||||
)
|
||||
|
||||
# 5a. Remove `custom` field from tools (Bedrock doesn't support it)
|
||||
# 5a. Bedrock Invoke supports output_config (effort) for Claude 4.6+ models,
|
||||
# but older models do not — strip it to avoid request rejection.
|
||||
# Ref: https://github.com/BerriAI/litellm/issues/22797
|
||||
if not (
|
||||
_supports_factory(
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
key="supports_output_config",
|
||||
)
|
||||
or AnthropicConfig._model_supports_effort_param(model)
|
||||
):
|
||||
if anthropic_messages_request.pop("output_config", None) is not None:
|
||||
verbose_logger.warning(
|
||||
"Bedrock Invoke: stripping unsupported `output_config` for "
|
||||
"model=%s — neither `supports_output_config` nor any "
|
||||
"`supports_*_reasoning_effort` flag is set in "
|
||||
"model_prices_and_context_window.json. Add the capability "
|
||||
"flag to the model JSON entry if this model accepts "
|
||||
"`output_config`.",
|
||||
model,
|
||||
)
|
||||
|
||||
# 5b. Remove `custom` field from tools (Bedrock doesn't support it)
|
||||
# Claude Code sends `custom: {defer_loading: true}` on tool definitions,
|
||||
# which causes Bedrock to reject the request with "Extra inputs are not permitted"
|
||||
# Ref: https://github.com/BerriAI/litellm/issues/22847
|
||||
|
|
|
|||
|
|
@ -16,7 +16,6 @@ from litellm.secret_managers.main import get_secret_str
|
|||
|
||||
from ...openai_like.chat.transformation import OpenAILikeChatConfig
|
||||
|
||||
|
||||
BEDROCK_MANTLE_DEFAULT_REGION = "us-east-1"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,5 @@
|
|||
import json
|
||||
from typing import Any, Optional
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from litellm.constants import STREAM_SSE_DONE_STRING
|
||||
from litellm.exceptions import AuthenticationError
|
||||
from litellm.litellm_core_utils.core_helpers import process_response_headers
|
||||
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
|
||||
|
|
@ -9,13 +7,17 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
|
|||
)
|
||||
from litellm.llms.openai.common_utils import OpenAIError
|
||||
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
|
||||
from litellm.responses.sse_output_recovery import (
|
||||
parse_sse_json_chunk,
|
||||
record_output_item_chunk,
|
||||
record_output_text_chunk,
|
||||
)
|
||||
from litellm.types.llms.openai import (
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.utils import CustomStreamWrapper
|
||||
|
||||
from ..authenticator import Authenticator
|
||||
from ..common_utils import (
|
||||
|
|
@ -111,86 +113,139 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
raw_response: Any,
|
||||
logging_obj: Any,
|
||||
):
|
||||
content_type = (raw_response.headers or {}).get("content-type", "")
|
||||
body_text = raw_response.text or ""
|
||||
if "text/event-stream" not in content_type.lower():
|
||||
trimmed_body = body_text.lstrip()
|
||||
if not (
|
||||
trimmed_body.startswith("event:")
|
||||
or trimmed_body.startswith("data:")
|
||||
or "\nevent:" in body_text
|
||||
or "\ndata:" in body_text
|
||||
):
|
||||
return super().transform_response_api_response(
|
||||
model=model,
|
||||
raw_response=raw_response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
if not self._should_parse_as_sse(
|
||||
raw_response=raw_response, body_text=body_text
|
||||
):
|
||||
return super().transform_response_api_response(
|
||||
model=model,
|
||||
raw_response=raw_response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
logging_obj.post_call(
|
||||
original_response=raw_response.text,
|
||||
additional_args={"complete_input_dict": {}},
|
||||
)
|
||||
|
||||
completed_response = None
|
||||
error_message = None
|
||||
for chunk in body_text.splitlines():
|
||||
stripped_chunk = CustomStreamWrapper._strip_sse_data_from_chunk(chunk)
|
||||
if not stripped_chunk:
|
||||
continue
|
||||
stripped_chunk = stripped_chunk.strip()
|
||||
if not stripped_chunk:
|
||||
continue
|
||||
if stripped_chunk == STREAM_SSE_DONE_STRING:
|
||||
break
|
||||
try:
|
||||
parsed_chunk = json.loads(stripped_chunk)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
if not isinstance(parsed_chunk, dict):
|
||||
continue
|
||||
event_type = parsed_chunk.get("type")
|
||||
if event_type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED:
|
||||
response_payload = parsed_chunk.get("response")
|
||||
if isinstance(response_payload, dict):
|
||||
response_payload = dict(response_payload)
|
||||
if "created_at" in response_payload:
|
||||
response_payload["created_at"] = _safe_convert_created_field(
|
||||
response_payload["created_at"]
|
||||
)
|
||||
try:
|
||||
completed_response = ResponsesAPIResponse(**response_payload)
|
||||
except Exception:
|
||||
completed_response = ResponsesAPIResponse.model_construct(
|
||||
**response_payload
|
||||
)
|
||||
break
|
||||
if event_type in (
|
||||
ResponsesAPIStreamEvents.RESPONSE_FAILED,
|
||||
ResponsesAPIStreamEvents.ERROR,
|
||||
):
|
||||
error_obj = parsed_chunk.get("error") or (
|
||||
parsed_chunk.get("response") or {}
|
||||
).get("error")
|
||||
if error_obj is not None:
|
||||
if isinstance(error_obj, dict):
|
||||
error_message = error_obj.get("message") or str(error_obj)
|
||||
else:
|
||||
error_message = str(error_obj)
|
||||
|
||||
completed_response, error_message = self._extract_completed_response_from_sse(
|
||||
body_text=body_text
|
||||
)
|
||||
if completed_response is None:
|
||||
raise OpenAIError(
|
||||
message=error_message or raw_response.text,
|
||||
status_code=raw_response.status_code,
|
||||
)
|
||||
|
||||
self._attach_response_headers(
|
||||
completed_response=completed_response, raw_response=raw_response
|
||||
)
|
||||
return completed_response
|
||||
|
||||
def _should_parse_as_sse(self, raw_response: Any, body_text: str) -> bool:
|
||||
content_type = (raw_response.headers or {}).get("content-type", "")
|
||||
if "text/event-stream" in content_type.lower():
|
||||
return True
|
||||
trimmed_body = body_text.lstrip()
|
||||
return bool(
|
||||
trimmed_body.startswith("event:")
|
||||
or trimmed_body.startswith("data:")
|
||||
or "\nevent:" in body_text
|
||||
or "\ndata:" in body_text
|
||||
)
|
||||
|
||||
def _extract_completed_response_from_sse(
|
||||
self, body_text: str
|
||||
) -> tuple[Optional[ResponsesAPIResponse], Optional[str]]:
|
||||
completed_response = None
|
||||
error_message = None
|
||||
streamed_output_items: Dict[int, dict] = {}
|
||||
text_only_output_items: Dict[int, dict] = {}
|
||||
for chunk in body_text.splitlines():
|
||||
parsed_chunk = parse_sse_json_chunk(chunk)
|
||||
if parsed_chunk is None:
|
||||
continue
|
||||
|
||||
event_type = parsed_chunk.get("type")
|
||||
if event_type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE:
|
||||
record_output_item_chunk(
|
||||
parsed_chunk=parsed_chunk,
|
||||
output_items=streamed_output_items,
|
||||
)
|
||||
continue
|
||||
|
||||
if event_type == ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE:
|
||||
record_output_text_chunk(
|
||||
parsed_chunk=parsed_chunk,
|
||||
output_items=streamed_output_items,
|
||||
text_only_items=text_only_output_items,
|
||||
)
|
||||
continue
|
||||
|
||||
if event_type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED:
|
||||
# Real OUTPUT_ITEM_DONE events take precedence at any given
|
||||
# output_index, but text-only items at indices without a
|
||||
# matching OUTPUT_ITEM_DONE must still be preserved (e.g.
|
||||
# providers that emit only OUTPUT_TEXT_DONE for some indices).
|
||||
merged_items: Dict[int, dict] = {**text_only_output_items}
|
||||
merged_items.update(streamed_output_items)
|
||||
completed_response = self._build_completed_response_from_chunk(
|
||||
parsed_chunk=parsed_chunk,
|
||||
streamed_output_items=merged_items,
|
||||
)
|
||||
break
|
||||
|
||||
if event_type in (
|
||||
ResponsesAPIStreamEvents.RESPONSE_FAILED,
|
||||
ResponsesAPIStreamEvents.ERROR,
|
||||
):
|
||||
extracted_error = self._extract_error_message(parsed_chunk)
|
||||
if extracted_error is not None:
|
||||
error_message = extracted_error
|
||||
|
||||
return completed_response, error_message
|
||||
|
||||
def _build_completed_response_from_chunk(
|
||||
self, parsed_chunk: Dict[str, Any], streamed_output_items: Dict[int, dict]
|
||||
) -> Optional[ResponsesAPIResponse]:
|
||||
response_payload = parsed_chunk.get("response")
|
||||
if not isinstance(response_payload, dict):
|
||||
return None
|
||||
response_payload = dict(response_payload)
|
||||
if not response_payload.get("output") and streamed_output_items:
|
||||
response_payload["output"] = [
|
||||
item for _, item in sorted(streamed_output_items.items())
|
||||
]
|
||||
if "created_at" in response_payload:
|
||||
response_payload["created_at"] = _safe_convert_created_field(
|
||||
response_payload["created_at"]
|
||||
)
|
||||
try:
|
||||
return ResponsesAPIResponse(**response_payload)
|
||||
except Exception:
|
||||
return ResponsesAPIResponse.model_construct(**response_payload)
|
||||
|
||||
def _extract_error_message(self, parsed_chunk: Dict[str, Any]) -> Optional[str]:
|
||||
error_obj = parsed_chunk.get("error") or (
|
||||
parsed_chunk.get("response") or {}
|
||||
).get("error")
|
||||
if error_obj is None:
|
||||
return None
|
||||
if isinstance(error_obj, dict):
|
||||
return error_obj.get("message") or str(error_obj)
|
||||
return str(error_obj)
|
||||
|
||||
def _attach_response_headers(
|
||||
self,
|
||||
completed_response: ResponsesAPIResponse,
|
||||
raw_response: Any,
|
||||
) -> None:
|
||||
raw_headers = dict(raw_response.headers)
|
||||
processed_headers = process_response_headers(raw_headers)
|
||||
if not hasattr(completed_response, "_hidden_params"):
|
||||
setattr(completed_response, "_hidden_params", {})
|
||||
completed_response._hidden_params["additional_headers"] = processed_headers
|
||||
completed_response._hidden_params["headers"] = raw_headers
|
||||
return completed_response
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Legacy /v1/embedding handler for Bedrock Cohere.
|
||||
Legacy /v1/embedding handler for Bedrock Cohere.
|
||||
"""
|
||||
|
||||
import json
|
||||
|
|
|
|||
|
|
@ -110,15 +110,35 @@ class CohereEmbeddingConfig:
|
|||
additional_args={"complete_input_dict": data},
|
||||
original_response=response_json,
|
||||
)
|
||||
return self._populate_embedding_response(
|
||||
response_json=response_json,
|
||||
model_response=model_response,
|
||||
model=model,
|
||||
encoding=encoding,
|
||||
input=input,
|
||||
)
|
||||
|
||||
def _populate_embedding_response(
|
||||
self,
|
||||
response_json: dict,
|
||||
model_response: EmbeddingResponse,
|
||||
model: str,
|
||||
encoding: Any,
|
||||
input: list,
|
||||
) -> EmbeddingResponse:
|
||||
"""
|
||||
response
|
||||
Parse a Cohere embed response body into an OpenAI-style EmbeddingResponse.
|
||||
|
||||
Split out from `_transform_response` so callers that already log
|
||||
`post_call` themselves (e.g. SageMaker's embedding handler) can reuse
|
||||
the parsing without triggering a second `post_call`.
|
||||
|
||||
Response shape:
|
||||
{
|
||||
'object': "list",
|
||||
'data': [
|
||||
|
||||
]
|
||||
'model',
|
||||
'usage'
|
||||
'data': [...],
|
||||
'model',
|
||||
'usage',
|
||||
}
|
||||
"""
|
||||
embeddings = response_json["embeddings"]
|
||||
|
|
@ -149,9 +169,6 @@ class CohereEmbeddingConfig:
|
|||
model_response.object = "list"
|
||||
model_response.data = output_data
|
||||
model_response.model = model
|
||||
input_tokens = 0
|
||||
for text in input:
|
||||
input_tokens += len(encoding.encode(text))
|
||||
|
||||
setattr(
|
||||
model_response,
|
||||
|
|
|
|||
|
|
@ -890,6 +890,18 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
)
|
||||
|
||||
# Some providers (e.g. OCI) require request signing after the body is built.
|
||||
# The default BaseConfig.sign_request returns (headers, None) — a no-op for
|
||||
# providers that don't need signing.
|
||||
headers, signed_body = provider_config.sign_request(
|
||||
headers=headers,
|
||||
optional_params=optional_params,
|
||||
request_data=data,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=input,
|
||||
|
|
@ -916,6 +928,7 @@ class BaseLLMHTTPHandler:
|
|||
client=client,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
signed_body=signed_body,
|
||||
)
|
||||
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
|
|
@ -926,12 +939,20 @@ class BaseLLMHTTPHandler:
|
|||
sync_httpx_client = client
|
||||
|
||||
try:
|
||||
response = sync_httpx_client.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
data=json.dumps(data),
|
||||
timeout=timeout,
|
||||
)
|
||||
if signed_body is not None:
|
||||
response = sync_httpx_client.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
data=signed_body,
|
||||
timeout=timeout,
|
||||
)
|
||||
else:
|
||||
response = sync_httpx_client.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
data=json.dumps(data),
|
||||
timeout=timeout,
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(
|
||||
e=e,
|
||||
|
|
@ -964,6 +985,7 @@ class BaseLLMHTTPHandler:
|
|||
api_key: Optional[str] = None,
|
||||
timeout: Optional[Union[float, httpx.Timeout]] = None,
|
||||
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
|
||||
signed_body: Optional[bytes] = None,
|
||||
) -> EmbeddingResponse:
|
||||
if client is None or not isinstance(client, AsyncHTTPHandler):
|
||||
async_httpx_client = get_async_httpx_client(
|
||||
|
|
@ -974,12 +996,20 @@ class BaseLLMHTTPHandler:
|
|||
async_httpx_client = client
|
||||
|
||||
try:
|
||||
response = await async_httpx_client.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
json=request_data,
|
||||
timeout=timeout,
|
||||
)
|
||||
if signed_body is not None:
|
||||
response = await async_httpx_client.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
data=signed_body,
|
||||
timeout=timeout,
|
||||
)
|
||||
else:
|
||||
response = await async_httpx_client.post(
|
||||
url=api_base,
|
||||
headers=headers,
|
||||
json=request_data,
|
||||
timeout=timeout,
|
||||
)
|
||||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=provider_config)
|
||||
|
||||
|
|
@ -1177,6 +1207,8 @@ class BaseLLMHTTPHandler:
|
|||
|
||||
data = transformed_result.data
|
||||
files = transformed_result.files
|
||||
if transformed_result.content_type is not None:
|
||||
headers["Content-Type"] = transformed_result.content_type
|
||||
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
|
|
@ -1409,6 +1441,8 @@ class BaseLLMHTTPHandler:
|
|||
document=document,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
# All providers return OCRRequestData
|
||||
|
|
@ -1477,6 +1511,8 @@ class BaseLLMHTTPHandler:
|
|||
document=document,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
# All providers return OCRRequestData
|
||||
|
|
@ -1852,7 +1888,9 @@ class BaseLLMHTTPHandler:
|
|||
async_httpx_client: AsyncHTTPHandler,
|
||||
request_url: str,
|
||||
headers: dict,
|
||||
signed_json_body: Optional[bytes],
|
||||
# str when the caller passes a pre-serialized (unsigned) body to avoid
|
||||
# re-dumping; bytes when a provider signed the request (e.g. Bedrock).
|
||||
signed_json_body: Optional[Union[str, bytes]],
|
||||
request_body: dict,
|
||||
stream: bool,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
|
|
@ -2043,8 +2081,18 @@ class BaseLLMHTTPHandler:
|
|||
model=model,
|
||||
)
|
||||
|
||||
# The request body was serialized once for the pre-call log input and
|
||||
# again for the wire (json.dumps is O(payload), large for long-context
|
||||
# Claude Code history). Serialize once and reuse for both. Only when
|
||||
# the provider didn't sign the request (sign_request no-op for the
|
||||
# native anthropic path -> signed_json_body is None); signed providers
|
||||
# (e.g. Bedrock) keep their signed body untouched. The HTTP-error
|
||||
# retry path mutates + re-signs the body, so it still re-serializes
|
||||
# internally -- this only deduplicates the success path.
|
||||
request_body_json = json.dumps(request_body)
|
||||
|
||||
logging_obj.pre_call(
|
||||
input=[{"role": "user", "content": json.dumps(request_body)}],
|
||||
input=[{"role": "user", "content": request_body_json}],
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": request_body,
|
||||
|
|
@ -2057,7 +2105,9 @@ class BaseLLMHTTPHandler:
|
|||
async_httpx_client=async_httpx_client,
|
||||
request_url=request_url,
|
||||
headers=headers,
|
||||
signed_json_body=signed_json_body,
|
||||
signed_json_body=(
|
||||
signed_json_body if signed_json_body is not None else request_body_json
|
||||
),
|
||||
request_body=request_body,
|
||||
stream=stream or False,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -2079,6 +2129,14 @@ class BaseLLMHTTPHandler:
|
|||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
if not self._has_agentic_completion_hook(logging_obj):
|
||||
# No callback overrides async_should_run_agentic_loop, so the
|
||||
# agentic wrapper's only effect would be buffering every chunk
|
||||
# and rebuilding the response from SSE at end-of-stream to call
|
||||
# hooks that all return (False, {}). Stream through directly and
|
||||
# skip that per-chunk + end-of-stream overhead.
|
||||
return completion_stream
|
||||
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
|
||||
AgenticAnthropicStreamingIterator,
|
||||
)
|
||||
|
|
@ -4586,6 +4644,51 @@ class BaseLLMHTTPHandler:
|
|||
fingerprints = list(kwargs.get("_agentic_loop_fingerprints", []) or [])
|
||||
return depth, max(max_loops, 1), fingerprints
|
||||
|
||||
@staticmethod
|
||||
def _has_agentic_completion_hook(logging_obj: Any) -> bool:
|
||||
"""
|
||||
True if any registered callback actually overrides
|
||||
``async_should_run_agentic_loop`` (the gate every agentic hook goes
|
||||
through). The base ``CustomLogger`` implementation returns
|
||||
``(False, {})``, so when nothing overrides it the agentic
|
||||
post-processing is a guaranteed no-op and the streaming wrapper that
|
||||
buffers + rebuilds the whole response from SSE just to call it can be
|
||||
skipped entirely.
|
||||
|
||||
Function-identity comparison (not a leaf ``__dict__`` check) so an
|
||||
override inherited through any intermediate class is still detected --
|
||||
a false negative here would silently disable agentic features.
|
||||
|
||||
String entries in ``litellm.callbacks`` (e.g. ``"datadog"``) are
|
||||
resolved to their ``CustomLogger`` instance via
|
||||
``get_custom_logger_compatible_class`` -- same pattern as
|
||||
``ProxyLogging._callback_capabilities`` -- so a string-registered
|
||||
agentic callback is detected too.
|
||||
"""
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
get_custom_logger_compatible_class,
|
||||
)
|
||||
|
||||
base_func = CustomLogger.async_should_run_agentic_loop
|
||||
callbacks = litellm.callbacks + (
|
||||
getattr(logging_obj, "dynamic_success_callbacks", None) or []
|
||||
)
|
||||
for cb in callbacks:
|
||||
if isinstance(cb, str):
|
||||
resolved = get_custom_logger_compatible_class(cb) # type: ignore[arg-type]
|
||||
if resolved is None:
|
||||
continue
|
||||
cb = resolved
|
||||
if not isinstance(cb, CustomLogger):
|
||||
continue
|
||||
cb_func = getattr(type(cb), "async_should_run_agentic_loop", base_func)
|
||||
if getattr(cb_func, "__func__", cb_func) is not getattr(
|
||||
base_func, "__func__", base_func
|
||||
):
|
||||
return True
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _check_agentic_loop_safety(
|
||||
tool_calls: Any,
|
||||
|
|
|
|||
|
|
@ -13,7 +13,6 @@ from typing import Tuple
|
|||
|
||||
import httpx
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pre-built response templates
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Cost calculator for Dashscope Chat models.
|
||||
Cost calculator for Dashscope Chat models.
|
||||
|
||||
Handles tiered pricing and prompt caching scenarios.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Support for OpenAI's `/v1/chat/completions` endpoint.
|
||||
Support for OpenAI's `/v1/chat/completions` endpoint.
|
||||
|
||||
Calls done in OpenAI/openai.py as DataRobot is openai-compatible.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Translate between Cohere's `/rerank` format and Deepinfra's `/rerank` format.
|
||||
Translate between Cohere's `/rerank` format and Deepinfra's `/rerank` format.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Cost calculator for DeepSeek Chat models.
|
||||
Cost calculator for DeepSeek Chat models.
|
||||
|
||||
Handles prompt caching scenario.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -23,7 +23,6 @@ from litellm.types.agents import (
|
|||
AgentVersionsResponse,
|
||||
)
|
||||
|
||||
|
||||
# Keys inside litellm_params that should be forwarded to the Gemini
|
||||
# create-agent body verbatim.
|
||||
_GEMINI_AGENT_BODY_KEYS = ("base_agent", "instructions", "base_environment")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
386
litellm/llms/oci/chat/cohere.py
Normal file
386
litellm/llms/oci/chat/cohere.py
Normal file
|
|
@ -0,0 +1,386 @@
|
|||
"""
|
||||
OCI Generative AI — Cohere-specific chat transformation helpers.
|
||||
|
||||
Handles message history building, tool definition adaptation, non-streaming
|
||||
response parsing, and streaming chunk parsing for models served with
|
||||
``apiFormat="COHERE"`` (e.g. ``cohere.command-*``).
|
||||
"""
|
||||
|
||||
import datetime
|
||||
import json
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import httpx
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.llms.oci.chat.generic import (
|
||||
_normalize_oci_finish_reason,
|
||||
_synthesize_oci_tool_call_id,
|
||||
)
|
||||
from litellm.llms.oci.common_utils import (
|
||||
OCI_JSON_TO_PYTHON_TYPES,
|
||||
OCIError,
|
||||
enrich_cohere_param_description,
|
||||
resolve_oci_schema_anyof,
|
||||
resolve_oci_schema_refs,
|
||||
sanitize_oci_schema,
|
||||
)
|
||||
from litellm.types.llms.oci import (
|
||||
CohereChatResult,
|
||||
CohereMessage,
|
||||
CohereParameterDefinition,
|
||||
CohereStreamChunk,
|
||||
CohereTool,
|
||||
CohereToolCall,
|
||||
CohereToolMessage,
|
||||
CohereToolResult,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import (
|
||||
Choices,
|
||||
Delta,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
StreamingChoices,
|
||||
)
|
||||
from litellm.types.utils import Usage
|
||||
|
||||
|
||||
def _extract_text_content(content: Any) -> str:
|
||||
"""Return the plain-text representation of a message content value."""
|
||||
if content is None:
|
||||
return ""
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
return "".join(
|
||||
item.get("text", "")
|
||||
for item in content
|
||||
if isinstance(item, dict) and item.get("type") == "text"
|
||||
)
|
||||
return str(content)
|
||||
|
||||
|
||||
def adapt_messages_to_cohere_standard(
|
||||
messages: List[AllMessageValues],
|
||||
) -> List[CohereMessage]:
|
||||
"""Build a Cohere ``chatHistory`` list from an OpenAI-format message array.
|
||||
|
||||
- All messages except the *last user message* are included. The caller pulls
|
||||
the last user message into the request's top-level ``message`` field, so
|
||||
trailing tool results (the standard agentic continuation pattern) still
|
||||
appear in ``chatHistory`` and reach the model.
|
||||
- If no user message exists, every message is included (no slice).
|
||||
- System messages must be filtered out by the caller (they are routed into
|
||||
``preambleOverride`` separately) — they are not represented in
|
||||
``chatHistory``.
|
||||
- Tool results are expressed as OCI ``CohereToolMessage.toolResults`` entries,
|
||||
with the originating call's name and parameters resolved from the preceding
|
||||
assistant message via a ``tool_call_id`` lookup.
|
||||
"""
|
||||
# First pass: build tool_call_id → CohereToolCall so tool-result messages can
|
||||
# reference the originating call by name and parameters.
|
||||
tool_call_lookup: Dict[str, CohereToolCall] = {}
|
||||
for msg in messages:
|
||||
if msg.get("role") == "assistant":
|
||||
tool_calls_raw: Any = msg.get("tool_calls") or []
|
||||
for tc in tool_calls_raw:
|
||||
tc_id = tc.get("id", "")
|
||||
raw_args: Any = tc.get("function", {}).get("arguments", "{}")
|
||||
try:
|
||||
params: Dict[str, Any] = (
|
||||
json.loads(raw_args) if isinstance(raw_args, str) else raw_args
|
||||
)
|
||||
except json.JSONDecodeError:
|
||||
params = {}
|
||||
tool_call_lookup[tc_id] = CohereToolCall(
|
||||
name=str(tc.get("function", {}).get("name", "")),
|
||||
parameters=params,
|
||||
)
|
||||
|
||||
last_user_index = next(
|
||||
(
|
||||
i
|
||||
for i in range(len(messages) - 1, -1, -1)
|
||||
if messages[i].get("role") == "user"
|
||||
),
|
||||
None,
|
||||
)
|
||||
history_source = (
|
||||
messages
|
||||
if last_user_index is None
|
||||
else [m for i, m in enumerate(messages) if i != last_user_index]
|
||||
)
|
||||
|
||||
chat_history: List[CohereMessage] = []
|
||||
for msg in history_source:
|
||||
role = msg.get("role")
|
||||
content = _extract_text_content(msg.get("content"))
|
||||
|
||||
tool_calls: Optional[List[CohereToolCall]] = None
|
||||
if role == "assistant" and msg.get("tool_calls"): # type: ignore[union-attr,typeddict-item]
|
||||
tool_calls = []
|
||||
for tc in msg["tool_calls"]: # type: ignore[union-attr,typeddict-item]
|
||||
raw_arguments: Any = tc.get("function", {}).get("arguments", {})
|
||||
if isinstance(raw_arguments, str):
|
||||
try:
|
||||
arguments: Dict[str, Any] = json.loads(raw_arguments)
|
||||
except json.JSONDecodeError:
|
||||
arguments = {}
|
||||
else:
|
||||
arguments = raw_arguments
|
||||
tool_calls.append(
|
||||
CohereToolCall(
|
||||
name=str(tc.get("function", {}).get("name", "")),
|
||||
parameters=arguments,
|
||||
)
|
||||
)
|
||||
|
||||
if role == "user":
|
||||
chat_history.append(CohereMessage(role="USER", message=content))
|
||||
elif role == "assistant":
|
||||
chat_history.append(
|
||||
CohereMessage(role="CHATBOT", message=content, toolCalls=tool_calls)
|
||||
)
|
||||
elif role == "tool":
|
||||
tool_call_id = str(msg.get("tool_call_id", "") or "")
|
||||
cohere_call = tool_call_lookup.get(
|
||||
tool_call_id, CohereToolCall(name="", parameters={})
|
||||
)
|
||||
tool_result = CohereToolResult(
|
||||
call=cohere_call,
|
||||
outputs=[{"output": content}],
|
||||
)
|
||||
# OpenAI emits one tool-role message per parallel tool call, but
|
||||
# the OCI Cohere API expects all results from a single assistant
|
||||
# turn to share one TOOL history entry with multiple toolResults.
|
||||
# Merge consecutive tool messages so the model sees the parallel
|
||||
# call/result pairing correctly during agentic loops.
|
||||
if chat_history and isinstance(chat_history[-1], CohereToolMessage):
|
||||
chat_history[-1].toolResults.append(tool_result)
|
||||
else:
|
||||
chat_history.append(CohereToolMessage(toolResults=[tool_result]))
|
||||
|
||||
return chat_history
|
||||
|
||||
|
||||
def adapt_tool_definitions_to_cohere_standard(
|
||||
tools: List[Dict[str, Any]],
|
||||
) -> List[CohereTool]:
|
||||
"""Adapt OpenAI-format tool definitions to the OCI Cohere format.
|
||||
|
||||
- Resolves ``$ref``/``$defs`` and ``anyOf`` patterns that OCI rejects.
|
||||
- Maps JSON Schema type names to Python type names (``"string"`` → ``"str"``).
|
||||
- Embeds unsupported constraints (enum, format, range, pattern) into the
|
||||
parameter description so the model can still see them.
|
||||
"""
|
||||
cohere_tools = []
|
||||
for tool in tools:
|
||||
function_def = tool.get("function", {})
|
||||
raw_params = function_def.get("parameters", {})
|
||||
|
||||
resolved = sanitize_oci_schema(
|
||||
resolve_oci_schema_anyof(resolve_oci_schema_refs(raw_params))
|
||||
)
|
||||
properties = resolved.get("properties", {})
|
||||
required = resolved.get("required", [])
|
||||
|
||||
parameter_definitions = {}
|
||||
for param_name, param_schema in properties.items():
|
||||
json_type = param_schema.get("type", "string")
|
||||
python_type = OCI_JSON_TO_PYTHON_TYPES.get(json_type, json_type)
|
||||
parameter_definitions[param_name] = CohereParameterDefinition(
|
||||
description=enrich_cohere_param_description(
|
||||
param_schema.get("description", ""), param_schema
|
||||
),
|
||||
type=python_type,
|
||||
isRequired=param_name in required,
|
||||
)
|
||||
|
||||
cohere_tools.append(
|
||||
CohereTool(
|
||||
name=function_def.get("name", ""),
|
||||
description=function_def.get("description", ""),
|
||||
parameterDefinitions=parameter_definitions,
|
||||
)
|
||||
)
|
||||
|
||||
return cohere_tools
|
||||
|
||||
|
||||
def handle_cohere_response(
|
||||
json_response: dict,
|
||||
model: str,
|
||||
model_response: ModelResponse,
|
||||
raw_response: httpx.Response,
|
||||
) -> ModelResponse:
|
||||
"""Parse a non-streaming Cohere OCI response into a LiteLLM ModelResponse."""
|
||||
try:
|
||||
cohere_response = CohereChatResult(**json_response)
|
||||
except (TypeError, ValidationError) as e:
|
||||
raise OCIError(
|
||||
message=f"Response cannot be casted to CohereChatResult: {str(e)}",
|
||||
status_code=raw_response.status_code,
|
||||
)
|
||||
|
||||
model_response.model = model
|
||||
model_response.created = int(datetime.datetime.now().timestamp())
|
||||
|
||||
response_text = cohere_response.chatResponse.text
|
||||
finish_reason = _normalize_oci_finish_reason(
|
||||
cohere_response.chatResponse.finishReason
|
||||
)
|
||||
|
||||
tool_calls: Optional[List[Dict[str, Any]]] = None
|
||||
if cohere_response.chatResponse.toolCalls:
|
||||
tool_calls = [
|
||||
{
|
||||
"id": _synthesize_oci_tool_call_id(
|
||||
i, tc.name, json.dumps(tc.parameters, sort_keys=True)
|
||||
),
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": tc.name,
|
||||
"arguments": json.dumps(tc.parameters),
|
||||
},
|
||||
}
|
||||
for i, tc in enumerate(cohere_response.chatResponse.toolCalls)
|
||||
]
|
||||
|
||||
content: Optional[str] = response_text if response_text else None
|
||||
|
||||
# Only include ``tool_calls`` in the message dict when actually present.
|
||||
# Passing an explicit ``None`` would let downstream consumers that key off
|
||||
# ``"tool_calls" in message`` (rather than truthiness) incorrectly conclude
|
||||
# that tool calls were attempted. Matches the generic handler's behaviour,
|
||||
# which only sets ``message.tool_calls`` when tool calls are present.
|
||||
message: Dict[str, Any] = {"role": "assistant", "content": content}
|
||||
if tool_calls is not None:
|
||||
message["tool_calls"] = tool_calls
|
||||
|
||||
model_response.choices = [
|
||||
Choices(
|
||||
index=0,
|
||||
message=message,
|
||||
finish_reason=finish_reason,
|
||||
)
|
||||
]
|
||||
|
||||
usage_info = cohere_response.chatResponse.usage
|
||||
if usage_info is not None:
|
||||
model_response.usage = Usage( # type: ignore[attr-defined]
|
||||
prompt_tokens=usage_info.promptTokens,
|
||||
completion_tokens=usage_info.completionTokens,
|
||||
total_tokens=usage_info.totalTokens,
|
||||
)
|
||||
else:
|
||||
model_response.usage = Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0) # type: ignore[attr-defined]
|
||||
|
||||
return model_response
|
||||
|
||||
|
||||
def handle_cohere_stream_chunk(
|
||||
dict_chunk: dict,
|
||||
prior_tool_calls_emitted: bool = False,
|
||||
prior_text_emitted: bool = False,
|
||||
) -> ModelResponseStream:
|
||||
"""Parse a single Cohere SSE chunk into a LiteLLM ModelResponseStream.
|
||||
|
||||
``prior_tool_calls_emitted`` lets the caller signal whether tool calls
|
||||
were already emitted in earlier chunks of the same stream. When set, the
|
||||
terminal consolidation chunk's tool calls are suppressed (they would
|
||||
duplicate prior deltas); otherwise they are passed through so a stream
|
||||
that delivers tool calls only on the terminal chunk doesn't silently
|
||||
drop them.
|
||||
|
||||
``prior_text_emitted`` plays the analogous role for the ``text`` field:
|
||||
when set, the terminal consolidation chunk's ``text`` is suppressed
|
||||
(it would re-emit the full assembled response on top of prior deltas);
|
||||
when unset (e.g. a degenerate stream that delivers the entire response
|
||||
in a single SSE event carrying both ``chatHistory`` and ``finishReason``),
|
||||
the text is passed through so the response content isn't silently lost.
|
||||
"""
|
||||
try:
|
||||
typed_chunk = CohereStreamChunk(**dict_chunk)
|
||||
except (TypeError, ValidationError) as e:
|
||||
raise OCIError(
|
||||
status_code=500,
|
||||
message=f"Chunk cannot be parsed as CohereStreamChunk: {str(e)}",
|
||||
)
|
||||
|
||||
if typed_chunk.index is None:
|
||||
typed_chunk.index = 0
|
||||
|
||||
# OCI Cohere's terminal SSE event re-sends the full assembled response in
|
||||
# `text` alongside a populated `chatHistory` and a non-null `finishReason`.
|
||||
# Emitting that text would concatenate the whole response onto the
|
||||
# already-streamed deltas. We require both signals to be present so that a
|
||||
# future API change which adds `chatHistory` to intermediate chunks (or a
|
||||
# rare early-populated case) doesn't silently drop legitimate token deltas.
|
||||
is_terminal_consolidation = (
|
||||
typed_chunk.chatHistory is not None and typed_chunk.finishReason is not None
|
||||
)
|
||||
# On non-terminal text-free chunks (e.g. tool-call-only or keep-alive
|
||||
# chunks) emit ``content=None`` rather than ``content=""`` so downstream
|
||||
# stream-mergers that distinguish "no text in this delta" from "an
|
||||
# explicitly empty text delta" behave correctly.
|
||||
#
|
||||
# We only suppress the terminal chunk's ``text`` when the caller has
|
||||
# confirmed that text deltas were already emitted earlier — otherwise
|
||||
# (e.g. a degenerate stream that delivers the whole response in a
|
||||
# single SSE event), passing it through is the only chance to surface it.
|
||||
text: Optional[str] = (
|
||||
None if (is_terminal_consolidation and prior_text_emitted) else typed_chunk.text
|
||||
)
|
||||
|
||||
# Tool calls on the terminal consolidation chunk (whether from
|
||||
# `typed_chunk.toolCalls` or from `chatHistory`) typically restate what
|
||||
# was already streamed in intermediate chunks. Re-emitting them would
|
||||
# mint fresh `uuid4` IDs and cause downstream consumers to execute each
|
||||
# tool call twice. We only suppress when the caller has confirmed that
|
||||
# tool calls were already emitted earlier — otherwise (e.g. a short
|
||||
# response that delivers tool calls exclusively on the terminal chunk),
|
||||
# passing them through is the only chance to surface them.
|
||||
cohere_tool_calls = (
|
||||
None
|
||||
if (is_terminal_consolidation and prior_tool_calls_emitted)
|
||||
else typed_chunk.toolCalls
|
||||
)
|
||||
|
||||
tool_calls: Optional[List[Dict[str, Any]]] = None
|
||||
if cohere_tool_calls:
|
||||
tool_calls = [
|
||||
{
|
||||
# Cohere protocol has no tool-call id, so we synthesize one
|
||||
# deterministically from the call's content/position. A random
|
||||
# uuid4 per chunk would cause downstream stream-mergers to
|
||||
# treat each chunk as a distinct tool call.
|
||||
"id": _synthesize_oci_tool_call_id(
|
||||
i, tc.name, json.dumps(tc.parameters, sort_keys=True)
|
||||
),
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": tc.name,
|
||||
"arguments": json.dumps(tc.parameters),
|
||||
},
|
||||
}
|
||||
for i, tc in enumerate(cohere_tool_calls)
|
||||
]
|
||||
|
||||
finish_reason = _normalize_oci_finish_reason(typed_chunk.finishReason)
|
||||
|
||||
return ModelResponseStream(
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=typed_chunk.index,
|
||||
delta=Delta(
|
||||
content=text,
|
||||
tool_calls=tool_calls,
|
||||
provider_specific_fields=None,
|
||||
thinking_blocks=None,
|
||||
reasoning_content=None,
|
||||
),
|
||||
finish_reason=finish_reason,
|
||||
)
|
||||
]
|
||||
)
|
||||
477
litellm/llms/oci/chat/generic.py
Normal file
477
litellm/llms/oci/chat/generic.py
Normal file
|
|
@ -0,0 +1,477 @@
|
|||
"""
|
||||
OCI Generative AI — Generic-format chat transformation helpers.
|
||||
|
||||
Handles message building, tool definition adaptation, non-streaming response
|
||||
parsing, and streaming chunk parsing for models served with
|
||||
``apiFormat="GENERIC"`` (e.g. Meta Llama, xAI Grok, Google Gemini).
|
||||
"""
|
||||
|
||||
import datetime
|
||||
import hashlib
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
import httpx
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.llms.oci.common_utils import (
|
||||
OCIError,
|
||||
resolve_oci_schema_anyof,
|
||||
resolve_oci_schema_refs,
|
||||
sanitize_oci_schema,
|
||||
)
|
||||
from litellm.types.llms.oci import (
|
||||
OCICompletionResponse,
|
||||
OCIContentPartUnion,
|
||||
OCIImageContentPart,
|
||||
OCIImageUrl,
|
||||
OCIMessage,
|
||||
OCIRoles,
|
||||
OCIStreamChunk,
|
||||
OCITextContentPart,
|
||||
OCIToolCall,
|
||||
OCIToolDefinition,
|
||||
OCIVendors,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import (
|
||||
Delta,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
StreamingChoices,
|
||||
)
|
||||
from litellm.types.utils import ChatCompletionMessageToolCall, Usage
|
||||
|
||||
# Maps OpenAI role names to OCI GENERIC role names.
|
||||
open_ai_to_generic_oci_role_map: Dict[str, OCIRoles] = {
|
||||
"system": "SYSTEM",
|
||||
"user": "USER",
|
||||
"assistant": "ASSISTANT",
|
||||
"tool": "TOOL",
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Message building
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def adapt_messages_to_generic_oci_standard_content_message(
|
||||
role: str, content: Union[str, list]
|
||||
) -> OCIMessage:
|
||||
"""Convert a plain-text or multipart content message to OCI format."""
|
||||
new_content: List[OCIContentPartUnion] = []
|
||||
if isinstance(content, str):
|
||||
return OCIMessage(
|
||||
role=open_ai_to_generic_oci_role_map[role],
|
||||
content=[OCITextContentPart(text=content)],
|
||||
toolCalls=None,
|
||||
toolCallId=None,
|
||||
)
|
||||
|
||||
for content_item in content:
|
||||
if not isinstance(content_item, dict):
|
||||
raise OCIError(
|
||||
status_code=400, message="Each content item must be a dictionary"
|
||||
)
|
||||
|
||||
item_type = content_item.get("type")
|
||||
if not isinstance(item_type, str):
|
||||
raise OCIError(
|
||||
status_code=400,
|
||||
message="Each content item must have a string `type` field",
|
||||
)
|
||||
if item_type not in ["text", "image_url"]:
|
||||
raise OCIError(
|
||||
status_code=400,
|
||||
message=f"Content type `{item_type}` is not supported by OCI",
|
||||
)
|
||||
|
||||
if item_type == "text":
|
||||
text = content_item.get("text")
|
||||
if not isinstance(text, str):
|
||||
raise OCIError(
|
||||
status_code=400,
|
||||
message="Content item of type `text` must have a string `text` field",
|
||||
)
|
||||
new_content.append(OCITextContentPart(text=text))
|
||||
|
||||
elif item_type == "image_url":
|
||||
image_url = content_item.get("image_url")
|
||||
if isinstance(image_url, dict):
|
||||
image_url = image_url.get("url")
|
||||
if not isinstance(image_url, str):
|
||||
raise OCIError(
|
||||
status_code=400,
|
||||
message="Prop `image_url` must be a string or an object with a `url` property",
|
||||
)
|
||||
new_content.append(OCIImageContentPart(imageUrl=OCIImageUrl(url=image_url)))
|
||||
|
||||
return OCIMessage(
|
||||
role=open_ai_to_generic_oci_role_map[role],
|
||||
content=new_content,
|
||||
toolCalls=None,
|
||||
toolCallId=None,
|
||||
)
|
||||
|
||||
|
||||
def adapt_messages_to_generic_oci_standard_tool_call(
|
||||
role: str, tool_calls: list
|
||||
) -> OCIMessage:
|
||||
"""Convert an assistant tool-call message to OCI format."""
|
||||
tool_calls_formatted = []
|
||||
for tool_call in tool_calls:
|
||||
if not isinstance(tool_call, dict):
|
||||
raise OCIError(
|
||||
status_code=400, message="Each tool call must be a dictionary"
|
||||
)
|
||||
if tool_call.get("type") != "function":
|
||||
raise OCIError(
|
||||
status_code=400, message="OCI only supports function tool calls"
|
||||
)
|
||||
|
||||
tool_call_id = tool_call.get("id")
|
||||
if not isinstance(tool_call_id, str):
|
||||
raise OCIError(status_code=400, message="Tool call `id` must be a string")
|
||||
|
||||
tool_function = tool_call.get("function")
|
||||
if not isinstance(tool_function, dict):
|
||||
raise OCIError(
|
||||
status_code=400, message="Tool call `function` must be a dictionary"
|
||||
)
|
||||
|
||||
function_name = tool_function.get("name")
|
||||
if not isinstance(function_name, str):
|
||||
raise OCIError(
|
||||
status_code=400, message="Tool call `function.name` must be a string"
|
||||
)
|
||||
|
||||
arguments = tool_call["function"].get("arguments", "{}")
|
||||
if not isinstance(arguments, str):
|
||||
raise OCIError(
|
||||
status_code=400,
|
||||
message="Tool call `function.arguments` must be a JSON string",
|
||||
)
|
||||
|
||||
tool_calls_formatted.append(
|
||||
OCIToolCall(
|
||||
id=tool_call_id,
|
||||
type="FUNCTION",
|
||||
name=function_name,
|
||||
arguments=arguments,
|
||||
)
|
||||
)
|
||||
|
||||
return OCIMessage(
|
||||
role=open_ai_to_generic_oci_role_map[role],
|
||||
content=None,
|
||||
toolCalls=tool_calls_formatted,
|
||||
toolCallId=None,
|
||||
)
|
||||
|
||||
|
||||
def adapt_messages_to_generic_oci_standard_tool_response(
|
||||
role: str, tool_call_id: str, content: str
|
||||
) -> OCIMessage:
|
||||
"""Convert a tool-result message to OCI format."""
|
||||
return OCIMessage(
|
||||
role=open_ai_to_generic_oci_role_map[role],
|
||||
content=[OCITextContentPart(text=content)],
|
||||
toolCalls=None,
|
||||
toolCallId=tool_call_id,
|
||||
)
|
||||
|
||||
|
||||
def adapt_messages_to_generic_oci_standard(
|
||||
messages: List[AllMessageValues],
|
||||
) -> List[OCIMessage]:
|
||||
"""Convert an OpenAI-format message array to OCI GENERIC format."""
|
||||
new_messages = []
|
||||
for message in messages:
|
||||
role = message["role"]
|
||||
content = message.get("content")
|
||||
tool_calls = message.get("tool_calls")
|
||||
tool_call_id = message.get("tool_call_id")
|
||||
|
||||
if role == "assistant" and tool_calls is not None:
|
||||
if not isinstance(tool_calls, list):
|
||||
raise OCIError(
|
||||
status_code=400, message="Message `tool_calls` must be a list"
|
||||
)
|
||||
new_messages.append(
|
||||
adapt_messages_to_generic_oci_standard_tool_call(role, tool_calls)
|
||||
)
|
||||
|
||||
elif role in ["system", "user", "assistant"] and content is not None:
|
||||
if not isinstance(content, (str, list)):
|
||||
raise OCIError(
|
||||
status_code=400,
|
||||
message="Message `content` must be a string or list of content parts",
|
||||
)
|
||||
new_messages.append(
|
||||
adapt_messages_to_generic_oci_standard_content_message(role, content)
|
||||
)
|
||||
|
||||
elif role == "tool":
|
||||
if not isinstance(tool_call_id, str):
|
||||
raise OCIError(
|
||||
status_code=400,
|
||||
message="Tool result message must have a string `tool_call_id`",
|
||||
)
|
||||
if not isinstance(content, str):
|
||||
raise OCIError(
|
||||
status_code=400,
|
||||
message="Tool result message `content` must be a string",
|
||||
)
|
||||
new_messages.append(
|
||||
adapt_messages_to_generic_oci_standard_tool_response(
|
||||
role, tool_call_id, content
|
||||
)
|
||||
)
|
||||
|
||||
return new_messages
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tool definition adaptation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def adapt_tool_definition_to_oci_standard(
|
||||
tools: List[Dict], vendor: OCIVendors
|
||||
) -> List[OCIToolDefinition]:
|
||||
"""Convert OpenAI-format tool definitions to OCI GENERIC format.
|
||||
|
||||
Resolves ``$ref``/``$defs`` and ``anyOf`` that the OCI endpoint rejects.
|
||||
"""
|
||||
new_tools = []
|
||||
for tool in tools:
|
||||
if tool["type"] != "function":
|
||||
raise OCIError(status_code=400, message="OCI only supports function tools")
|
||||
|
||||
tool_function = tool.get("function")
|
||||
if not isinstance(tool_function, dict):
|
||||
raise OCIError(
|
||||
status_code=400, message="Tool `function` must be a dictionary"
|
||||
)
|
||||
|
||||
raw_params = tool_function.get("parameters", {})
|
||||
resolved_params = sanitize_oci_schema(
|
||||
resolve_oci_schema_anyof(resolve_oci_schema_refs(raw_params))
|
||||
)
|
||||
|
||||
new_tools.append(
|
||||
OCIToolDefinition(
|
||||
type="FUNCTION",
|
||||
name=tool_function.get("name"),
|
||||
description=tool_function.get("description", ""),
|
||||
parameters=resolved_params,
|
||||
)
|
||||
)
|
||||
|
||||
return new_tools
|
||||
|
||||
|
||||
def _normalize_oci_finish_reason(raw: Optional[str]) -> Optional[str]:
|
||||
"""Map an OCI-specific finish reason to its OpenAI-standard equivalent.
|
||||
|
||||
OCI emits ``COMPLETE`` / ``MAX_TOKENS`` / ``TOOL_CALL(S)`` plus a long tail
|
||||
of error/cancel reasons (``ERROR``, ``ERROR_TOXIC``, ``ERROR_LIMIT``,
|
||||
``USER_CANCEL``, ``CONTENT_FILTERED``, ``CANCELLED``, ...). The OpenAI
|
||||
spec only defines ``stop`` / ``length`` / ``tool_calls`` / ... — anything
|
||||
else is collapsed to ``"stop"`` so downstream consumers switching on
|
||||
``finish_reason`` keep working. A ``None`` input passes through unchanged.
|
||||
"""
|
||||
if raw is None:
|
||||
return None
|
||||
if raw == "COMPLETE":
|
||||
return "stop"
|
||||
if raw == "MAX_TOKENS":
|
||||
return "length"
|
||||
if raw in ("TOOL_CALL", "TOOL_CALLS"):
|
||||
return "tool_calls"
|
||||
return "stop"
|
||||
|
||||
|
||||
def _synthesize_oci_tool_call_id(position: int, name: str, arguments: str) -> str:
|
||||
"""Deterministic synthetic tool-call id derived from chunk content.
|
||||
|
||||
Used as a fallback when OCI omits ``id`` (always the case for the OCI
|
||||
Cohere protocol, occasionally the case for OCI GENERIC streaming chunks).
|
||||
A random ``uuid4`` per chunk would cause downstream stream-merging
|
||||
consumers — which key off the tool-call ``id`` — to treat re-emissions of
|
||||
the same logical call (e.g. terminal consolidation chunks, retries) as
|
||||
distinct calls. A content-derived digest stays stable across identical
|
||||
re-emissions while differing across truly distinct calls.
|
||||
"""
|
||||
digest = hashlib.sha256(
|
||||
f"{position}|{name}|{arguments}".encode("utf-8"),
|
||||
usedforsecurity=False,
|
||||
).hexdigest()[:24]
|
||||
return f"call_{digest}"
|
||||
|
||||
|
||||
def adapt_tools_to_openai_standard(
|
||||
tools: List[OCIToolCall],
|
||||
) -> List[ChatCompletionMessageToolCall]:
|
||||
"""Convert OCI tool-call objects in a response to the OpenAI format."""
|
||||
return [
|
||||
ChatCompletionMessageToolCall(
|
||||
id=tool.id or _synthesize_oci_tool_call_id(i, tool.name, tool.arguments),
|
||||
type="function",
|
||||
function={"name": tool.name, "arguments": tool.arguments},
|
||||
)
|
||||
for i, tool in enumerate(tools)
|
||||
]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Response parsing
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def handle_generic_response(
|
||||
json_data: dict,
|
||||
model: str,
|
||||
model_response: ModelResponse,
|
||||
raw_response: httpx.Response,
|
||||
) -> ModelResponse:
|
||||
"""Parse a non-streaming GENERIC OCI response into a LiteLLM ModelResponse."""
|
||||
try:
|
||||
completion_response = OCICompletionResponse(**json_data)
|
||||
except (TypeError, ValidationError) as e:
|
||||
raise OCIError(
|
||||
message=f"Response cannot be casted to OCICompletionResponse: {str(e)}",
|
||||
status_code=raw_response.status_code,
|
||||
)
|
||||
|
||||
iso_str = completion_response.chatResponse.timeCreated
|
||||
dt = datetime.datetime.fromisoformat(iso_str.replace("Z", "+00:00"))
|
||||
model_response.created = int(dt.timestamp())
|
||||
model_response.model = completion_response.modelId
|
||||
|
||||
if not completion_response.chatResponse.choices:
|
||||
raise OCIError(
|
||||
message="OCI response contained no choices",
|
||||
status_code=raw_response.status_code,
|
||||
)
|
||||
|
||||
response_choice = completion_response.chatResponse.choices[0]
|
||||
message = model_response.choices[0].message # type: ignore
|
||||
response_message = response_choice.message
|
||||
if response_message is not None:
|
||||
if response_message.content:
|
||||
# Concatenate all text parts — matches the streaming handler, which
|
||||
# iterates the full content array. Skips non-text parts (e.g. image
|
||||
# parts) so a leading non-text part doesn't suppress trailing text.
|
||||
text: Optional[str] = None
|
||||
for item in response_message.content:
|
||||
if isinstance(item, OCITextContentPart):
|
||||
text = (text or "") + item.text
|
||||
if text is not None:
|
||||
message.content = text
|
||||
if response_message.toolCalls:
|
||||
message.tool_calls = adapt_tools_to_openai_standard(
|
||||
response_message.toolCalls
|
||||
)
|
||||
|
||||
model_response.choices[0].finish_reason = _normalize_oci_finish_reason( # type: ignore[union-attr,assignment]
|
||||
response_choice.finishReason
|
||||
)
|
||||
|
||||
oci_usage = completion_response.chatResponse.usage
|
||||
reasoning_tokens: Optional[int] = None
|
||||
if (
|
||||
oci_usage.completionTokensDetails
|
||||
and oci_usage.completionTokensDetails.reasoningTokens is not None
|
||||
):
|
||||
reasoning_tokens = oci_usage.completionTokensDetails.reasoningTokens
|
||||
model_response.usage = Usage( # type: ignore[attr-defined]
|
||||
prompt_tokens=oci_usage.promptTokens,
|
||||
completion_tokens=oci_usage.completionTokens or 0,
|
||||
total_tokens=oci_usage.totalTokens,
|
||||
reasoning_tokens=reasoning_tokens,
|
||||
)
|
||||
|
||||
return model_response
|
||||
|
||||
|
||||
def handle_generic_stream_chunk(dict_chunk: dict) -> ModelResponseStream:
|
||||
"""Parse a single GENERIC SSE chunk into a LiteLLM ModelResponseStream."""
|
||||
# OCI streams tool calls progressively — early chunks may omit required fields.
|
||||
if dict_chunk.get("message") and dict_chunk["message"].get("toolCalls"):
|
||||
for tool_call in dict_chunk["message"]["toolCalls"]:
|
||||
tool_call.setdefault("arguments", "")
|
||||
tool_call.setdefault("id", "")
|
||||
tool_call.setdefault("name", "")
|
||||
|
||||
try:
|
||||
typed_chunk = OCIStreamChunk(**dict_chunk)
|
||||
except (TypeError, ValidationError) as e:
|
||||
raise OCIError(
|
||||
status_code=500,
|
||||
message=f"Chunk cannot be parsed as OCIStreamChunk: {str(e)}",
|
||||
)
|
||||
|
||||
if typed_chunk.index is None:
|
||||
typed_chunk.index = 0
|
||||
|
||||
# Emit ``content=None`` rather than ``content=""`` on chunks with no text
|
||||
# parts (e.g. tool-call-only or keep-alive chunks) so downstream
|
||||
# stream-mergers that distinguish "no text in this delta" from "an
|
||||
# explicitly empty text delta" behave correctly.
|
||||
text: Optional[str] = None
|
||||
if typed_chunk.message and typed_chunk.message.content:
|
||||
for item in typed_chunk.message.content:
|
||||
if isinstance(item, OCITextContentPart):
|
||||
text = (text or "") + item.text
|
||||
elif isinstance(item, OCIImageContentPart):
|
||||
raise OCIError(
|
||||
status_code=500,
|
||||
message="OCI returned image content in a streaming response — not supported",
|
||||
)
|
||||
else:
|
||||
raise OCIError(
|
||||
status_code=500,
|
||||
message=f"Unsupported content type in OCI streaming response: {item.type}",
|
||||
)
|
||||
|
||||
# Build plain tool-call dicts inline (matching the shape produced by
|
||||
# ``handle_cohere_stream_chunk``) rather than calling
|
||||
# ``adapt_tools_to_openai_standard`` and ``model_dump``-ing the typed
|
||||
# objects. Both code paths feed ``Delta.tool_calls``, so emitting the
|
||||
# same minimal ``{"id", "type", "function": {"name", "arguments"}}``
|
||||
# shape keeps downstream stream-mergers behaving identically across
|
||||
# GENERIC and Cohere chunks.
|
||||
tool_calls: Optional[List[Dict[str, Any]]] = None
|
||||
if typed_chunk.message and typed_chunk.message.toolCalls:
|
||||
tool_calls = [
|
||||
{
|
||||
"id": tc.id or _synthesize_oci_tool_call_id(i, tc.name, tc.arguments),
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": tc.name,
|
||||
"arguments": tc.arguments,
|
||||
},
|
||||
}
|
||||
for i, tc in enumerate(typed_chunk.message.toolCalls)
|
||||
]
|
||||
|
||||
finish_reason: Optional[str] = _normalize_oci_finish_reason(
|
||||
typed_chunk.finishReason
|
||||
)
|
||||
|
||||
return ModelResponseStream(
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=typed_chunk.index,
|
||||
delta=Delta(
|
||||
content=text,
|
||||
tool_calls=tool_calls,
|
||||
provider_specific_fields=None,
|
||||
thinking_blocks=None,
|
||||
reasoning_content=None,
|
||||
),
|
||||
finish_reason=finish_reason,
|
||||
)
|
||||
]
|
||||
)
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -1,9 +1,42 @@
|
|||
from typing import Optional
|
||||
import base64
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from email.utils import formatdate
|
||||
from typing import Any, Dict, Optional, Protocol, Tuple
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
|
||||
try:
|
||||
from cryptography.hazmat.primitives import hashes, serialization
|
||||
from cryptography.hazmat.primitives.asymmetric import padding, rsa
|
||||
|
||||
_CRYPTOGRAPHY_AVAILABLE = True
|
||||
except ImportError:
|
||||
_CRYPTOGRAPHY_AVAILABLE = False
|
||||
|
||||
try:
|
||||
from litellm._version import version as _litellm_version
|
||||
except ImportError:
|
||||
_litellm_version = "0.0.0"
|
||||
|
||||
|
||||
# OCI GenAI REST API version — stable since service launch, unlikely to change
|
||||
OCI_API_VERSION = "20231130"
|
||||
|
||||
|
||||
def _require_cryptography() -> None:
|
||||
if not _CRYPTOGRAPHY_AVAILABLE:
|
||||
raise ImportError(
|
||||
"cryptography package is required for OCI authentication. "
|
||||
"Please install it with: pip install cryptography"
|
||||
)
|
||||
|
||||
|
||||
class OCIError(BaseLLMException):
|
||||
def __init__(
|
||||
|
|
@ -17,3 +50,520 @@ class OCIError(BaseLLMException):
|
|||
message=message,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# OCI signing protocol and helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class OCISignerProtocol(Protocol):
|
||||
"""
|
||||
Protocol for OCI request signers (e.g., oci.signer.Signer).
|
||||
|
||||
Compatible with the OCI Python SDK's Signer class.
|
||||
See: https://docs.oracle.com/en-us/iaas/tools/python/latest/api/signing.html
|
||||
"""
|
||||
|
||||
def do_request_sign(
|
||||
self, request: Any, *, enforce_content_headers: bool = False
|
||||
) -> None:
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class OCIRequestWrapper:
|
||||
"""
|
||||
Wrapper for HTTP requests compatible with OCI signer interface.
|
||||
|
||||
Wraps request data in the format expected by OCI SDK signers, which require
|
||||
objects with method, url, headers, body, and path_url attributes.
|
||||
"""
|
||||
|
||||
method: str
|
||||
url: str
|
||||
headers: dict
|
||||
body: bytes
|
||||
|
||||
@property
|
||||
def path_url(self) -> str:
|
||||
"""Returns the path + query string for OCI signing."""
|
||||
parsed = urlparse(self.url)
|
||||
return parsed.path + ("?" + parsed.query if parsed.query else "")
|
||||
|
||||
|
||||
def sha256_base64(data: bytes) -> str:
|
||||
# SHA-256 is used here to compute the x-content-sha256 header required by the
|
||||
# OCI HTTP signing specification (RSA-SHA256 request signing), not for password
|
||||
# or secret hashing. This is the correct and mandated algorithm for this purpose.
|
||||
# See: https://docs.oracle.com/en-us/iaas/Content/API/Concepts/signingrequests.htm
|
||||
#
|
||||
# ``usedforsecurity=False`` declares non-security intent to static analyzers
|
||||
# (CodeQL ``py/weak-sensitive-data-hashing``) — without it the request body
|
||||
# gets flagged as "password-like data" via taint tracking.
|
||||
digest = hashlib.sha256(data, usedforsecurity=False).digest() # noqa: S324
|
||||
return base64.b64encode(digest).decode()
|
||||
|
||||
|
||||
def build_signature_string(
|
||||
method: str, path: str, headers: dict, signed_headers: list
|
||||
) -> str:
|
||||
lines = []
|
||||
for header in signed_headers:
|
||||
if header == "(request-target)":
|
||||
value = f"{method.lower()} {path}"
|
||||
else:
|
||||
value = headers[header]
|
||||
lines.append(f"{header}: {value}")
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def load_private_key_from_str(key_str: str) -> Any:
|
||||
_require_cryptography()
|
||||
key = serialization.load_pem_private_key( # type: ignore[union-attr]
|
||||
key_str.encode("utf-8"),
|
||||
password=None,
|
||||
)
|
||||
if not isinstance(key, rsa.RSAPrivateKey): # type: ignore[union-attr]
|
||||
raise TypeError(
|
||||
"The provided private key is not an RSA key, which is required for OCI signing."
|
||||
)
|
||||
return key
|
||||
|
||||
|
||||
def load_private_key_from_file(file_path: str) -> Any:
|
||||
"""Loads a private key from a file path."""
|
||||
try:
|
||||
with open(file_path, "r", encoding="utf-8") as f:
|
||||
key_str = f.read().strip()
|
||||
except FileNotFoundError:
|
||||
raise FileNotFoundError(f"Private key file not found: {file_path}")
|
||||
except OSError as e:
|
||||
raise OSError(f"Failed to read private key file '{file_path}': {e}") from e
|
||||
|
||||
if not key_str:
|
||||
raise ValueError(f"Private key file is empty: {file_path}")
|
||||
|
||||
return load_private_key_from_str(key_str)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Env-var credential resolution
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_OCI_REGION_ENV = "OCI_REGION"
|
||||
_OCI_USER_ENV = "OCI_USER"
|
||||
_OCI_FINGERPRINT_ENV = "OCI_FINGERPRINT"
|
||||
_OCI_TENANCY_ENV = "OCI_TENANCY"
|
||||
_OCI_KEY_FILE_ENV = "OCI_KEY_FILE"
|
||||
_OCI_KEY_ENV = "OCI_KEY"
|
||||
_OCI_COMPARTMENT_ID_ENV = "OCI_COMPARTMENT_ID"
|
||||
|
||||
|
||||
def resolve_oci_credentials(optional_params: dict) -> dict:
|
||||
"""
|
||||
Merge OCI credentials from optional_params (explicit, always wins) and
|
||||
environment variables (fallback).
|
||||
|
||||
Returns a dict with resolved values for:
|
||||
oci_region, oci_user, oci_fingerprint, oci_tenancy,
|
||||
oci_key, oci_key_file, oci_compartment_id
|
||||
"""
|
||||
return {
|
||||
"oci_region": optional_params.get("oci_region")
|
||||
or os.environ.get(_OCI_REGION_ENV)
|
||||
or "us-ashburn-1",
|
||||
"oci_user": optional_params.get("oci_user") or os.environ.get(_OCI_USER_ENV),
|
||||
"oci_fingerprint": optional_params.get("oci_fingerprint")
|
||||
or os.environ.get(_OCI_FINGERPRINT_ENV),
|
||||
"oci_tenancy": optional_params.get("oci_tenancy")
|
||||
or os.environ.get(_OCI_TENANCY_ENV),
|
||||
"oci_key": optional_params.get("oci_key") or os.environ.get(_OCI_KEY_ENV),
|
||||
"oci_key_file": optional_params.get("oci_key_file")
|
||||
or os.environ.get(_OCI_KEY_FILE_ENV),
|
||||
"oci_compartment_id": optional_params.get("oci_compartment_id")
|
||||
or os.environ.get(_OCI_COMPARTMENT_ID_ENV),
|
||||
}
|
||||
|
||||
|
||||
_OCI_REGION_RE = re.compile(r"^[a-z][a-z0-9-]{0,30}[a-z0-9]$")
|
||||
_OCI_ACTION_PATH_RE = re.compile(rf"/{OCI_API_VERSION}/actions/[^/?#]+/?$")
|
||||
|
||||
|
||||
def get_oci_base_url(optional_params: dict, api_base: Optional[str] = None) -> str:
|
||||
"""Return the OCI inference base URL, respecting any explicit api_base override.
|
||||
|
||||
If ``api_base`` already ends with a fully-formed OCI action path
|
||||
(``/{OCI_API_VERSION}/actions/<name>``), that suffix is stripped so callers
|
||||
can append their own action path without producing a doubled URL.
|
||||
"""
|
||||
if api_base:
|
||||
return _OCI_ACTION_PATH_RE.sub("", api_base).rstrip("/")
|
||||
creds = resolve_oci_credentials(optional_params)
|
||||
region = creds["oci_region"]
|
||||
if not isinstance(region, str) or not _OCI_REGION_RE.match(region):
|
||||
raise OCIError(
|
||||
status_code=400,
|
||||
message=(
|
||||
f"Invalid OCI region {region!r}: must match "
|
||||
"^[a-z][a-z0-9-]{0,30}[a-z0-9]$ (e.g. 'us-ashburn-1')."
|
||||
),
|
||||
)
|
||||
return f"https://inference.generativeai.{region}.oci.oraclecloud.com"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Signing implementations (shared by chat, embed, and rerank configs)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def sign_with_oci_signer(
|
||||
headers: dict,
|
||||
optional_params: dict,
|
||||
request_data: dict,
|
||||
api_base: str,
|
||||
) -> Tuple[dict, bytes]:
|
||||
"""Sign a request using an OCI SDK Signer object passed in optional_params."""
|
||||
oci_signer = optional_params.get("oci_signer")
|
||||
body = json.dumps(request_data).encode("utf-8")
|
||||
method = str(optional_params.get("method", "POST")).upper()
|
||||
|
||||
if method not in {"POST", "GET", "PUT", "DELETE", "PATCH"}:
|
||||
raise ValueError(f"Unsupported HTTP method: {method}")
|
||||
|
||||
prepared_headers = {**headers}
|
||||
prepared_headers.setdefault("content-type", "application/json")
|
||||
prepared_headers.setdefault("content-length", str(len(body)))
|
||||
|
||||
request_wrapper = OCIRequestWrapper(
|
||||
method=method, url=api_base, headers=prepared_headers, body=body
|
||||
)
|
||||
|
||||
if oci_signer is None:
|
||||
raise ValueError("oci_signer cannot be None when calling sign_with_oci_signer")
|
||||
|
||||
try:
|
||||
oci_signer.do_request_sign(request_wrapper, enforce_content_headers=True)
|
||||
except Exception as e:
|
||||
raise OCIError(
|
||||
status_code=500,
|
||||
message=(
|
||||
f"Failed to sign request with provided oci_signer: {str(e)}. "
|
||||
"The signer must implement the OCI SDK Signer interface with a "
|
||||
"do_request_sign(request, enforce_content_headers=True) method. "
|
||||
"See: https://docs.oracle.com/en-us/iaas/tools/python/latest/api/signing.html"
|
||||
),
|
||||
) from e
|
||||
|
||||
headers.update(request_wrapper.headers)
|
||||
return headers, body
|
||||
|
||||
|
||||
def sign_with_manual_credentials(
|
||||
headers: dict,
|
||||
optional_params: dict,
|
||||
request_data: dict,
|
||||
api_base: str,
|
||||
) -> Tuple[dict, bytes]:
|
||||
"""Sign a request using manually provided OCI credentials (user/fingerprint/tenancy/key)."""
|
||||
creds = resolve_oci_credentials(optional_params)
|
||||
oci_user = creds["oci_user"]
|
||||
oci_fingerprint = creds["oci_fingerprint"]
|
||||
oci_tenancy = creds["oci_tenancy"]
|
||||
oci_key = creds["oci_key"]
|
||||
oci_key_file = creds["oci_key_file"]
|
||||
|
||||
if (
|
||||
not oci_user
|
||||
or not oci_fingerprint
|
||||
or not oci_tenancy
|
||||
or not (oci_key or oci_key_file)
|
||||
):
|
||||
raise OCIError(
|
||||
status_code=401,
|
||||
message=(
|
||||
"Missing required OCI credentials: oci_user, oci_fingerprint, oci_tenancy, "
|
||||
"and at least one of oci_key or oci_key_file. "
|
||||
"These can also be supplied via environment variables: "
|
||||
f"{_OCI_USER_ENV}, {_OCI_FINGERPRINT_ENV}, {_OCI_TENANCY_ENV}, {_OCI_KEY_ENV} (or {_OCI_KEY_FILE_ENV}). "
|
||||
"Alternatively, provide an oci_signer object from the OCI SDK."
|
||||
),
|
||||
)
|
||||
|
||||
method = str(optional_params.get("method", "POST")).upper()
|
||||
body = json.dumps(request_data).encode("utf-8")
|
||||
parsed = urlparse(api_base)
|
||||
path = parsed.path or "/"
|
||||
host = parsed.netloc
|
||||
|
||||
date = formatdate(usegmt=True)
|
||||
content_type = headers.get("content-type", "application/json")
|
||||
content_length = str(len(body))
|
||||
x_content_sha256 = sha256_base64(body)
|
||||
|
||||
headers_to_sign: Dict[str, str] = {
|
||||
"date": date,
|
||||
"host": host,
|
||||
"content-type": content_type,
|
||||
"content-length": content_length,
|
||||
"x-content-sha256": x_content_sha256,
|
||||
}
|
||||
|
||||
signed_header_names = [
|
||||
"date",
|
||||
"(request-target)",
|
||||
"host",
|
||||
"content-length",
|
||||
"content-type",
|
||||
"x-content-sha256",
|
||||
]
|
||||
signing_string = build_signature_string(
|
||||
method, path, headers_to_sign, signed_header_names
|
||||
)
|
||||
|
||||
_require_cryptography()
|
||||
|
||||
# Resolve the private key — prefer inline PEM content over file path
|
||||
oci_key_content: Optional[str] = None
|
||||
if oci_key:
|
||||
if not isinstance(oci_key, str):
|
||||
raise OCIError(
|
||||
status_code=400,
|
||||
message=(
|
||||
f"oci_key must be a string containing the PEM private key content. "
|
||||
f"Got type: {type(oci_key).__name__}"
|
||||
),
|
||||
)
|
||||
oci_key_content = oci_key.replace("\\n", "\n").replace("\r\n", "\n")
|
||||
|
||||
private_key = (
|
||||
load_private_key_from_str(oci_key_content)
|
||||
if oci_key_content
|
||||
else load_private_key_from_file(oci_key_file) if oci_key_file else None
|
||||
)
|
||||
|
||||
if private_key is None:
|
||||
raise OCIError(
|
||||
status_code=400,
|
||||
message="Private key is required for OCI authentication. Provide either oci_key or oci_key_file.",
|
||||
)
|
||||
|
||||
signature = private_key.sign(
|
||||
signing_string.encode("utf-8"),
|
||||
padding.PKCS1v15(), # type: ignore[union-attr]
|
||||
hashes.SHA256(), # type: ignore[union-attr]
|
||||
)
|
||||
signature_b64 = base64.b64encode(signature).decode()
|
||||
|
||||
key_id = f"{oci_tenancy}/{oci_user}/{oci_fingerprint}"
|
||||
authorization = (
|
||||
'Signature version="1",'
|
||||
f'keyId="{key_id}",'
|
||||
'algorithm="rsa-sha256",'
|
||||
f'headers="{" ".join(signed_header_names)}",'
|
||||
f'signature="{signature_b64}"'
|
||||
)
|
||||
|
||||
headers.update(
|
||||
{
|
||||
"authorization": authorization,
|
||||
"date": date,
|
||||
"host": host,
|
||||
"content-type": content_type,
|
||||
"content-length": content_length,
|
||||
"x-content-sha256": x_content_sha256,
|
||||
}
|
||||
)
|
||||
return headers, body
|
||||
|
||||
|
||||
def sign_oci_request(
|
||||
headers: dict,
|
||||
optional_params: dict,
|
||||
request_data: dict,
|
||||
api_base: str,
|
||||
api_key: Optional[str] = None,
|
||||
model: Optional[str] = None,
|
||||
stream: Optional[bool] = None,
|
||||
fake_stream: Optional[bool] = None,
|
||||
) -> Tuple[dict, bytes]:
|
||||
"""
|
||||
Route to the appropriate OCI signing method based on what credentials are present.
|
||||
|
||||
If ``oci_signer`` is in optional_params, use the OCI SDK signer object.
|
||||
Otherwise use manual RSA-SHA256 signing with explicit credentials (which can
|
||||
also be supplied via OCI_* environment variables).
|
||||
|
||||
Returns:
|
||||
Tuple of (signed_headers, signed_body_bytes)
|
||||
"""
|
||||
if optional_params.get("oci_signer") is not None:
|
||||
return sign_with_oci_signer(headers, optional_params, request_data, api_base)
|
||||
return sign_with_manual_credentials(
|
||||
headers, optional_params, request_data, api_base
|
||||
)
|
||||
|
||||
|
||||
def validate_oci_environment(
|
||||
headers: dict,
|
||||
optional_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""
|
||||
Populate common OCI request headers (content-type, user-agent).
|
||||
|
||||
Full credential validation is deferred to signing time so that credentials
|
||||
supplied via environment variables are resolved at call time rather than
|
||||
at construction time.
|
||||
"""
|
||||
headers.setdefault("content-type", "application/json")
|
||||
headers.setdefault("user-agent", f"litellm/{_litellm_version}")
|
||||
return headers
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# JSON schema utilities for OCI tool definitions
|
||||
#
|
||||
# OCI Generative AI does not support JSON Schema extensions ($ref, $defs,
|
||||
# anyOf). Pydantic v2 emits all three for models with Optional fields or
|
||||
# nested schemas. The helpers below are ported from the official
|
||||
# langchain-oracle reference implementation so that tool schemas are always
|
||||
# valid before they reach the OCI endpoint.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Mapping from JSON Schema type names to Python type names, as expected by
|
||||
# the OCI Cohere API's CohereParameterDefinition.type field.
|
||||
OCI_JSON_TO_PYTHON_TYPES: Dict[str, str] = {
|
||||
"string": "str",
|
||||
"number": "float",
|
||||
"boolean": "bool",
|
||||
"integer": "int",
|
||||
"array": "List",
|
||||
"object": "Dict",
|
||||
"any": "any",
|
||||
}
|
||||
|
||||
|
||||
def resolve_oci_schema_refs(schema: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Inline all ``$ref``/``$defs`` references — OCI does not support JSON Schema ``$ref``."""
|
||||
defs = schema.get("$defs", {})
|
||||
resolving_stack: set = set()
|
||||
|
||||
def _resolve(obj: Any) -> Any:
|
||||
if isinstance(obj, dict):
|
||||
if "$ref" in obj:
|
||||
ref = obj["$ref"]
|
||||
if ref.startswith("#/$defs/"):
|
||||
key = ref.split("/")[-1]
|
||||
if key in resolving_stack:
|
||||
return {"type": "object"} # break cycles
|
||||
resolving_stack.add(key)
|
||||
try:
|
||||
return _resolve(defs.get(key, obj))
|
||||
finally:
|
||||
resolving_stack.discard(key)
|
||||
return obj # external $ref — leave unchanged
|
||||
return {k: _resolve(v) for k, v in obj.items()}
|
||||
if isinstance(obj, list):
|
||||
return [_resolve(item) for item in obj]
|
||||
return obj
|
||||
|
||||
resolved = _resolve(schema)
|
||||
if isinstance(resolved, dict):
|
||||
resolved.pop("$defs", None)
|
||||
return resolved
|
||||
|
||||
|
||||
def resolve_oci_schema_anyof(obj: Any) -> Any:
|
||||
"""Resolve Pydantic v2 ``Optional[T]`` → ``anyOf`` patterns.
|
||||
|
||||
Pydantic v2 emits ``{"anyOf": [{"type": "T"}, {"type": "null"}]}`` for
|
||||
``Optional[T]``. OCI models don't understand ``anyOf``, so we pick the
|
||||
first non-null branch and merge top-level metadata into it.
|
||||
"""
|
||||
if isinstance(obj, dict):
|
||||
if "anyOf" in obj and "type" not in obj:
|
||||
non_null = [
|
||||
t
|
||||
for t in obj["anyOf"]
|
||||
if not (isinstance(t, dict) and t.get("type") == "null")
|
||||
]
|
||||
if non_null:
|
||||
resolved = {**obj, **non_null[0]}
|
||||
resolved.pop("anyOf", None)
|
||||
return resolve_oci_schema_anyof(resolved)
|
||||
return {k: resolve_oci_schema_anyof(v) for k, v in obj.items()}
|
||||
if isinstance(obj, list):
|
||||
return [resolve_oci_schema_anyof(item) for item in obj]
|
||||
return obj
|
||||
|
||||
|
||||
def sanitize_oci_schema(schema: Any) -> Any:
|
||||
"""Recursively remove OCI-incompatible fields from a JSON schema.
|
||||
|
||||
Strips ``title`` keys, removes ``None``-valued ``default`` entries,
|
||||
normalises ``type: [T, "null"]`` list types, and ensures arrays carry an
|
||||
``items`` definition.
|
||||
"""
|
||||
if isinstance(schema, list):
|
||||
return [sanitize_oci_schema(item) for item in schema]
|
||||
if not isinstance(schema, dict):
|
||||
return schema
|
||||
|
||||
sanitized: Dict[str, Any] = {}
|
||||
for key, value in schema.items():
|
||||
if key == "title":
|
||||
continue
|
||||
if key == "default" and value is None:
|
||||
continue
|
||||
if key == "type":
|
||||
if value == "any":
|
||||
sanitized[key] = "object"
|
||||
continue
|
||||
if isinstance(value, list):
|
||||
non_null = [t for t in value if t != "null"]
|
||||
sanitized[key] = non_null[0] if non_null else "string"
|
||||
continue
|
||||
sanitized[key] = sanitize_oci_schema(value)
|
||||
|
||||
if sanitized.get("type") == "array" and "items" not in sanitized:
|
||||
sanitized["items"] = {"type": "object"}
|
||||
|
||||
required = sanitized.get("required")
|
||||
properties = sanitized.get("properties")
|
||||
if "required" in sanitized:
|
||||
if isinstance(required, list) and isinstance(properties, dict):
|
||||
sanitized["required"] = [
|
||||
f for f in required if isinstance(f, str) and f in properties
|
||||
]
|
||||
elif not isinstance(required, list):
|
||||
sanitized["required"] = []
|
||||
|
||||
return sanitized
|
||||
|
||||
|
||||
def enrich_cohere_param_description(
|
||||
description: str, param_schema: Dict[str, Any]
|
||||
) -> str:
|
||||
"""Embed schema constraints into a Cohere parameter description.
|
||||
|
||||
``CohereParameterDefinition`` only has ``type``, ``description``, and
|
||||
``isRequired``. Rich constraints (``enum``, ``format``, ``minimum``,
|
||||
``maximum``, ``pattern``) are appended to the description string so the
|
||||
model can still see and respect them.
|
||||
"""
|
||||
parts = [description] if description else []
|
||||
if "enum" in param_schema:
|
||||
parts.append(f"Allowed values: {param_schema['enum']}")
|
||||
if "format" in param_schema:
|
||||
parts.append(f"Format: {param_schema['format']}")
|
||||
if "minimum" in param_schema or "maximum" in param_schema:
|
||||
range_parts = []
|
||||
if "minimum" in param_schema:
|
||||
range_parts.append(f"min={param_schema['minimum']}")
|
||||
if "maximum" in param_schema:
|
||||
range_parts.append(f"max={param_schema['maximum']}")
|
||||
parts.append(f"Range: {', '.join(range_parts)}")
|
||||
if "pattern" in param_schema:
|
||||
parts.append(f"Pattern: {param_schema['pattern']}")
|
||||
return ". ".join(parts) if parts else ""
|
||||
|
|
|
|||
|
|
@ -1,8 +1,14 @@
|
|||
"""
|
||||
OCI Generative AI Embedding Configuration
|
||||
OCI Generative AI — Embedding transformation.
|
||||
|
||||
Supports embedding models available on Oracle Cloud Infrastructure Generative AI service.
|
||||
Uses the same authentication mechanisms as OCI chat (manual signing or OCI SDK Signer).
|
||||
Endpoint: POST /20231130/actions/embedText
|
||||
Supported models: cohere.embed-english-v3.0, cohere.embed-multilingual-v3.0,
|
||||
cohere.embed-v4.0, and all other Cohere embed variants available on OCI
|
||||
(including dedicated endpoints).
|
||||
|
||||
Authentication follows the same RSA-SHA256 / OCI SDK signer pattern as chat.
|
||||
The base handler (base_llm_http_handler.embedding) calls sign_request after
|
||||
building the body, so signing happens automatically.
|
||||
|
||||
Supported models:
|
||||
- cohere.embed-english-v3.0
|
||||
|
|
@ -10,25 +16,45 @@ Supported models:
|
|||
- cohere.embed-multilingual-v3.0
|
||||
- cohere.embed-multilingual-light-v3.0
|
||||
- cohere.embed-english-image-v3.0
|
||||
- cohere.embed-english-light-image-v3.0
|
||||
- cohere.embed-multilingual-light-image-v3.0
|
||||
- cohere.embed-multilingual-image-v3.0
|
||||
- cohere.embed-v4.0
|
||||
|
||||
Reference: https://docs.oracle.com/en-us/iaas/api/#/en/generative-ai-inference/latest/EmbedTextResult/EmbedText
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
import litellm
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
|
||||
from litellm.llms.oci.chat.transformation import OCIChatConfig
|
||||
from litellm.llms.oci.common_utils import OCIError
|
||||
from litellm.llms.oci.common_utils import (
|
||||
OCI_API_VERSION,
|
||||
OCIError,
|
||||
get_oci_base_url,
|
||||
resolve_oci_credentials,
|
||||
sign_oci_request,
|
||||
validate_oci_environment,
|
||||
)
|
||||
from litellm.types.llms.oci import (
|
||||
OCIEmbedRequest,
|
||||
OCIEmbedResponse,
|
||||
OCIServingMode,
|
||||
)
|
||||
from litellm.types.llms.openai import AllEmbeddingInputValues, AllMessageValues
|
||||
from litellm.types.utils import EmbeddingResponse, Usage
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
|
||||
# OCI sends up to 96 texts per embedText request (Cohere limit).
|
||||
OCI_EMBED_BATCH_LIMIT = 96
|
||||
|
||||
# Input type mapping from OpenAI conventions to OCI/Cohere conventions
|
||||
_INPUT_TYPE_MAP = {
|
||||
"search_document": "SEARCH_DOCUMENT",
|
||||
|
|
@ -38,65 +64,43 @@ _INPUT_TYPE_MAP = {
|
|||
}
|
||||
|
||||
|
||||
class OCIEmbeddingConfig(BaseEmbeddingConfig):
|
||||
class OCIEmbedConfig(BaseEmbeddingConfig):
|
||||
"""
|
||||
Configuration for OCI Generative AI Embedding API.
|
||||
Transformation config for OCI Generative AI embeddings.
|
||||
|
||||
The OCI embedding endpoint uses the Cohere embed models hosted on OCI.
|
||||
Authentication is handled via OCI request signing (manual credentials or OCI SDK Signer).
|
||||
Supports both text and (on cohere.embed-v4.0) multimodal inputs.
|
||||
|
||||
Usage:
|
||||
```python
|
||||
import litellm
|
||||
Authentication — same two modes as chat:
|
||||
- **OCI SDK signer**: pass ``oci_signer`` in optional_params.
|
||||
- **Manual RSA-SHA256**: pass ``oci_user``, ``oci_fingerprint``, ``oci_tenancy``,
|
||||
and ``oci_key`` or ``oci_key_file``, or set the corresponding ``OCI_*`` env vars.
|
||||
|
||||
response = litellm.embedding(
|
||||
model="oci/cohere.embed-english-v3.0",
|
||||
input=["Hello world", "Goodbye world"],
|
||||
oci_compartment_id="ocid1.compartment.oc1..xxx",
|
||||
oci_region="us-ashburn-1",
|
||||
oci_user="ocid1.user.oc1..xxx",
|
||||
oci_fingerprint="xx:xx:xx:xx",
|
||||
oci_tenancy="ocid1.tenancy.oc1..xxx",
|
||||
oci_key_file="~/.oci/key.pem",
|
||||
)
|
||||
```
|
||||
Required call-time params (via optional_params or env vars):
|
||||
- ``oci_compartment_id`` / ``OCI_COMPARTMENT_ID``
|
||||
- ``oci_region`` / ``OCI_REGION`` (default: ``us-ashburn-1``)
|
||||
|
||||
Optional call-time params:
|
||||
- ``oci_serving_mode``: ``"ON_DEMAND"`` (default) or ``"DEDICATED"``
|
||||
- ``oci_endpoint_id``: endpoint OCID for dedicated serving mode
|
||||
- ``input_type``: ``SEARCH_DOCUMENT``, ``SEARCH_QUERY``, ``CLASSIFICATION``, ``CLUSTERING``
|
||||
- ``truncate``: ``NONE``, ``START``, or ``END`` (default ``END``)
|
||||
- ``dimensions``: output embedding dimensions (cohere.embed-v4.0+)
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
# We reuse OCIChatConfig for signing logic
|
||||
self._chat_config = OCIChatConfig()
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
if api_base:
|
||||
return api_base
|
||||
|
||||
oci_region = optional_params.get("oci_region", "us-ashburn-1")
|
||||
return f"https://inference.generativeai.{oci_region}.oci.oraclecloud.com/20231130/actions/embedText"
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list:
|
||||
return [
|
||||
"dimensions",
|
||||
]
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
return ["dimensions"]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
drop_params: bool = False,
|
||||
) -> dict:
|
||||
# Note: OCI Cohere embed does not support custom dimensions natively,
|
||||
# but we pass it through in case future models support it
|
||||
if "dimensions" in non_default_params:
|
||||
optional_params["dimensions"] = non_default_params["dimensions"]
|
||||
for key, value in non_default_params.items():
|
||||
if key == "dimensions":
|
||||
# OCI API uses outputDimensions (cohere.embed-v4.0+)
|
||||
optional_params["outputDimensions"] = value
|
||||
return optional_params
|
||||
|
||||
def validate_environment(
|
||||
|
|
@ -109,49 +113,42 @@ class OCIEmbeddingConfig(BaseEmbeddingConfig):
|
|||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""
|
||||
Validate OCI credentials for embedding requests.
|
||||
Supports both OCI SDK Signer and manual credential signing.
|
||||
"""
|
||||
oci_signer = optional_params.get("oci_signer")
|
||||
oci_region = optional_params.get("oci_region", "us-ashburn-1")
|
||||
|
||||
api_base = (
|
||||
api_base
|
||||
or f"https://inference.generativeai.{oci_region}.oci.oraclecloud.com"
|
||||
)
|
||||
|
||||
if oci_signer is None:
|
||||
oci_user = optional_params.get("oci_user")
|
||||
oci_fingerprint = optional_params.get("oci_fingerprint")
|
||||
oci_tenancy = optional_params.get("oci_tenancy")
|
||||
oci_key = optional_params.get("oci_key")
|
||||
oci_key_file = optional_params.get("oci_key_file")
|
||||
oci_compartment_id = optional_params.get("oci_compartment_id")
|
||||
|
||||
if (
|
||||
not oci_user
|
||||
or not oci_fingerprint
|
||||
or not oci_tenancy
|
||||
or not (oci_key or oci_key_file)
|
||||
or not oci_compartment_id
|
||||
):
|
||||
raise Exception(
|
||||
"Missing required parameters: oci_user, oci_fingerprint, oci_tenancy, oci_compartment_id "
|
||||
"and at least one of oci_key or oci_key_file. "
|
||||
"Alternatively, provide an oci_signer object from the OCI SDK."
|
||||
if optional_params.get("oci_signer") is None:
|
||||
creds = resolve_oci_credentials(optional_params)
|
||||
missing = [
|
||||
k
|
||||
for k in (
|
||||
"oci_user",
|
||||
"oci_fingerprint",
|
||||
"oci_tenancy",
|
||||
"oci_compartment_id",
|
||||
)
|
||||
if not creds.get(k)
|
||||
]
|
||||
if missing or not (creds.get("oci_key") or creds.get("oci_key_file")):
|
||||
raise OCIError(
|
||||
status_code=401,
|
||||
message=(
|
||||
"Missing required parameters: oci_user, oci_fingerprint, oci_tenancy, "
|
||||
"oci_compartment_id and at least one of oci_key or oci_key_file. "
|
||||
"These can be supplied via optional_params or via OCI_USER, OCI_FINGERPRINT, "
|
||||
"OCI_TENANCY, OCI_COMPARTMENT_ID, OCI_KEY_FILE environment variables. "
|
||||
"Alternatively, provide an oci_signer object from the OCI SDK."
|
||||
),
|
||||
)
|
||||
return validate_oci_environment(headers, optional_params, api_key)
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import version
|
||||
|
||||
headers.update(
|
||||
{
|
||||
"content-type": "application/json",
|
||||
"user-agent": f"litellm/{version}",
|
||||
}
|
||||
)
|
||||
|
||||
return headers
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
base = get_oci_base_url(optional_params, api_base or litellm.api_base)
|
||||
return f"{base}/{OCI_API_VERSION}/actions/embedText"
|
||||
|
||||
def sign_request(
|
||||
self,
|
||||
|
|
@ -163,9 +160,8 @@ class OCIEmbeddingConfig(BaseEmbeddingConfig):
|
|||
model: Optional[str] = None,
|
||||
stream: Optional[bool] = None,
|
||||
fake_stream: Optional[bool] = None,
|
||||
):
|
||||
"""Delegate to OCIChatConfig's signing logic."""
|
||||
return self._chat_config.sign_request(
|
||||
) -> Tuple[dict, bytes]:
|
||||
return sign_oci_request(
|
||||
headers=headers,
|
||||
optional_params=optional_params,
|
||||
request_data=request_data,
|
||||
|
|
@ -182,91 +178,74 @@ class OCIEmbeddingConfig(BaseEmbeddingConfig):
|
|||
input: AllEmbeddingInputValues,
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform the embedding request to OCI format.
|
||||
|
||||
OCI embedText API expects:
|
||||
{
|
||||
"compartmentId": "...",
|
||||
"servingMode": {"servingType": "ON_DEMAND", "modelId": "..."},
|
||||
"inputs": ["text1", "text2"],
|
||||
"truncate": "END",
|
||||
"inputType": "SEARCH_DOCUMENT"
|
||||
}
|
||||
"""
|
||||
oci_compartment_id = optional_params.get("oci_compartment_id")
|
||||
if not oci_compartment_id:
|
||||
raise Exception(
|
||||
"kwarg `oci_compartment_id` is required for OCI embedding requests"
|
||||
creds = resolve_oci_credentials(optional_params)
|
||||
compartment_id = creds["oci_compartment_id"]
|
||||
if not compartment_id:
|
||||
raise OCIError(
|
||||
status_code=400,
|
||||
message=(
|
||||
"oci_compartment_id is required for OCI embedding requests. "
|
||||
"Pass it as optional_params or set the OCI_COMPARTMENT_ID env var."
|
||||
),
|
||||
)
|
||||
|
||||
# Build serving mode
|
||||
oci_serving_mode = optional_params.get("oci_serving_mode", "ON_DEMAND")
|
||||
if oci_serving_mode == "DEDICATED":
|
||||
oci_endpoint_id = optional_params.get("oci_endpoint_id", model)
|
||||
serving_mode = {
|
||||
"servingType": "DEDICATED",
|
||||
"endpointId": oci_endpoint_id,
|
||||
}
|
||||
else:
|
||||
serving_mode = {
|
||||
"servingType": "ON_DEMAND",
|
||||
"modelId": model,
|
||||
}
|
||||
|
||||
# Normalize input to list of strings
|
||||
# Normalise input to a flat list of strings
|
||||
if isinstance(input, str):
|
||||
inputs = [input]
|
||||
texts = [input]
|
||||
elif isinstance(input, list):
|
||||
inputs = []
|
||||
texts = []
|
||||
for item in input:
|
||||
if isinstance(item, str):
|
||||
inputs.append(item)
|
||||
elif isinstance(item, list):
|
||||
raise ValueError(
|
||||
"OCI embedding does not support token-array inputs. "
|
||||
"Please convert token lists to strings before calling embedding()."
|
||||
if isinstance(item, list):
|
||||
raise OCIError(
|
||||
status_code=400,
|
||||
message=(
|
||||
"OCI embedText does not support token-array inputs. "
|
||||
"Convert token lists to strings before calling embedding()."
|
||||
),
|
||||
)
|
||||
else:
|
||||
inputs.append(str(item))
|
||||
texts.append(item if isinstance(item, str) else str(item))
|
||||
else:
|
||||
inputs = [str(input)]
|
||||
texts = [str(input)]
|
||||
|
||||
# Build request data — OCI embedText API expects inputs, truncate,
|
||||
# and inputType at the top level alongside compartmentId and servingMode
|
||||
request_data: Dict[str, Any] = {
|
||||
"compartmentId": oci_compartment_id,
|
||||
"servingMode": serving_mode,
|
||||
"inputs": inputs,
|
||||
"truncate": optional_params.get("truncate", "END"),
|
||||
}
|
||||
if len(texts) > OCI_EMBED_BATCH_LIMIT:
|
||||
raise OCIError(
|
||||
status_code=400,
|
||||
message=(
|
||||
f"OCI embedText accepts at most {OCI_EMBED_BATCH_LIMIT} inputs per request "
|
||||
f"(got {len(texts)}). Batch your requests."
|
||||
),
|
||||
)
|
||||
|
||||
# Map input_type if provided
|
||||
serving_mode_type = optional_params.get("oci_serving_mode", "ON_DEMAND").upper()
|
||||
if serving_mode_type not in {"ON_DEMAND", "DEDICATED"}:
|
||||
raise OCIError(
|
||||
status_code=400,
|
||||
message="oci_serving_mode must be 'ON_DEMAND' or 'DEDICATED'.",
|
||||
)
|
||||
|
||||
if serving_mode_type == "DEDICATED":
|
||||
endpoint_id = optional_params.get("oci_endpoint_id", model)
|
||||
serving_mode = OCIServingMode(
|
||||
servingType="DEDICATED", endpointId=endpoint_id
|
||||
)
|
||||
else:
|
||||
serving_mode = OCIServingMode(servingType="ON_DEMAND", modelId=model)
|
||||
|
||||
# Map input_type from OpenAI convention to OCI/Cohere convention
|
||||
input_type = optional_params.get("input_type")
|
||||
if input_type:
|
||||
mapped_type = _INPUT_TYPE_MAP.get(input_type.lower(), input_type.upper())
|
||||
request_data["inputType"] = mapped_type
|
||||
input_type = _INPUT_TYPE_MAP.get(input_type.lower(), input_type.upper())
|
||||
|
||||
# Sign the request using the same URL the HTTP handler will POST to
|
||||
signing_url = self.get_complete_url(
|
||||
api_base=api_base,
|
||||
api_key=None,
|
||||
model=model,
|
||||
optional_params=optional_params,
|
||||
litellm_params={},
|
||||
request = OCIEmbedRequest(
|
||||
compartmentId=compartment_id,
|
||||
servingMode=serving_mode,
|
||||
inputs=texts,
|
||||
inputType=input_type,
|
||||
truncate=optional_params.get("truncate", "END"),
|
||||
outputDimensions=optional_params.get("outputDimensions"),
|
||||
)
|
||||
|
||||
signed_headers, body = self.sign_request(
|
||||
headers=headers,
|
||||
optional_params=optional_params,
|
||||
request_data=request_data,
|
||||
api_base=signing_url,
|
||||
)
|
||||
headers.update(signed_headers)
|
||||
|
||||
return request_data
|
||||
return request.model_dump(exclude_none=True)
|
||||
|
||||
def transform_embedding_response(
|
||||
self,
|
||||
|
|
@ -274,63 +253,57 @@ class OCIEmbeddingConfig(BaseEmbeddingConfig):
|
|||
raw_response: httpx.Response,
|
||||
model_response: EmbeddingResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
api_key: Optional[str] = None,
|
||||
request_data: dict = {},
|
||||
optional_params: dict = {},
|
||||
litellm_params: dict = {},
|
||||
api_key: Optional[str],
|
||||
request_data: dict,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
) -> EmbeddingResponse:
|
||||
"""
|
||||
Transform OCI embedding response to standard EmbeddingResponse format.
|
||||
|
||||
OCI response format:
|
||||
{
|
||||
"embeddings": [[0.1, 0.2, ...], [0.3, 0.4, ...]],
|
||||
"modelId": "cohere.embed-english-v3.0",
|
||||
"modelVersion": "3.0",
|
||||
"inputTextTokenCounts": [5, 4]
|
||||
}
|
||||
"""
|
||||
if raw_response.status_code != 200:
|
||||
raise OCIError(
|
||||
message=raw_response.text,
|
||||
status_code=raw_response.status_code,
|
||||
message=raw_response.text,
|
||||
)
|
||||
|
||||
try:
|
||||
raw_response_json = raw_response.json()
|
||||
except Exception:
|
||||
json_response = raw_response.json()
|
||||
except Exception as e:
|
||||
raise OCIError(
|
||||
message=raw_response.text,
|
||||
status_code=raw_response.status_code,
|
||||
message=f"Failed to parse OCI embed response as JSON: {e}",
|
||||
)
|
||||
|
||||
embeddings = raw_response_json.get("embeddings", [])
|
||||
model_id = raw_response_json.get("modelId", model)
|
||||
|
||||
# Build response data in OpenAI format
|
||||
embedding_data = []
|
||||
for idx, embedding in enumerate(embeddings):
|
||||
embedding_data.append(
|
||||
{
|
||||
"object": "embedding",
|
||||
"index": idx,
|
||||
"embedding": embedding,
|
||||
}
|
||||
try:
|
||||
parsed = OCIEmbedResponse(**json_response)
|
||||
except Exception as e:
|
||||
raise OCIError(
|
||||
status_code=500,
|
||||
message=f"OCI embed response does not match expected schema: {e}",
|
||||
)
|
||||
|
||||
model_response.model = model_id
|
||||
model_response.data = embedding_data
|
||||
model_response.object = "list"
|
||||
model_response.model = parsed.modelId
|
||||
model_response.data = [
|
||||
{
|
||||
"object": "embedding",
|
||||
"index": i,
|
||||
"embedding": embedding,
|
||||
}
|
||||
for i, embedding in enumerate(parsed.embeddings)
|
||||
]
|
||||
|
||||
# Calculate token usage
|
||||
input_token_counts = raw_response_json.get("inputTextTokenCounts", [])
|
||||
total_tokens = sum(input_token_counts) if input_token_counts else 0
|
||||
|
||||
usage = Usage(
|
||||
prompt_tokens=total_tokens,
|
||||
total_tokens=total_tokens,
|
||||
)
|
||||
model_response.usage = usage
|
||||
if parsed.inputTextTokenCounts is not None:
|
||||
# Actual OCI API returns per-input token counts — sum for total usage
|
||||
total = sum(parsed.inputTextTokenCounts)
|
||||
model_response.usage = Usage(prompt_tokens=total, total_tokens=total)
|
||||
elif parsed.usage is not None:
|
||||
# Some deployments may return a usage object directly
|
||||
model_response.usage = Usage(
|
||||
prompt_tokens=parsed.usage.promptTokens,
|
||||
total_tokens=parsed.usage.totalTokens,
|
||||
)
|
||||
else:
|
||||
# Neither field returned — default to zero so downstream consumers
|
||||
# can always rely on usage being populated.
|
||||
model_response.usage = Usage(prompt_tokens=0, total_tokens=0)
|
||||
|
||||
return model_response
|
||||
|
||||
|
|
@ -340,8 +313,8 @@ class OCIEmbeddingConfig(BaseEmbeddingConfig):
|
|||
status_code: int,
|
||||
headers: Union[dict, httpx.Headers],
|
||||
) -> BaseLLMException:
|
||||
return OCIError(
|
||||
message=error_message,
|
||||
status_code=status_code,
|
||||
headers=headers if isinstance(headers, httpx.Headers) else None,
|
||||
)
|
||||
return OCIError(status_code=status_code, message=error_message)
|
||||
|
||||
|
||||
# Alias for backwards compatibility with any code that imports OCIEmbeddingConfig
|
||||
OCIEmbeddingConfig = OCIEmbedConfig
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -126,9 +126,21 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict,
|
||||
) -> Dict:
|
||||
"""No transform applied since inputs are in OpenAI spec already"""
|
||||
"""Strip Anthropic-only `cache_control` markers before sending to OpenAI.
|
||||
|
||||
OpenAI's Responses API rejects unknown fields on input content blocks
|
||||
with HTTP 400 ("Unknown parameter: 'input[0].content[0].cache_control'").
|
||||
Chat Completions strips these in
|
||||
`remove_cache_control_flag_from_messages_and_tools`; mirror that here.
|
||||
"""
|
||||
|
||||
input = self._validate_input_param(input)
|
||||
tools = response_api_optional_request_params.get("tools")
|
||||
input, tools = self.remove_cache_control_flag_from_input_and_tools(
|
||||
model=model, input=input, tools=tools
|
||||
)
|
||||
if tools is not None:
|
||||
response_api_optional_request_params["tools"] = tools
|
||||
final_request_params = dict(
|
||||
ResponsesAPIRequestParams(
|
||||
model=model, input=input, **response_api_optional_request_params
|
||||
|
|
@ -137,6 +149,38 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
|
||||
return final_request_params
|
||||
|
||||
def remove_cache_control_flag_from_input_and_tools(
|
||||
self,
|
||||
model: str, # allows overrides to selectively run this
|
||||
input: Union[str, ResponseInputParam],
|
||||
tools: Optional[List[ALL_RESPONSES_API_TOOL_PARAMS]] = None,
|
||||
) -> Tuple[
|
||||
Union[str, ResponseInputParam],
|
||||
Optional[List[ALL_RESPONSES_API_TOOL_PARAMS]],
|
||||
]:
|
||||
"""Sibling of `remove_cache_control_flag_from_messages_and_tools` on
|
||||
the chat path. Strips Anthropic-only `cache_control` markers from
|
||||
Responses API input content blocks and tools.
|
||||
|
||||
`filter_value_from_dict` mutates each dict in place, so the same
|
||||
objects are returned.
|
||||
"""
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
filter_value_from_dict,
|
||||
)
|
||||
|
||||
if isinstance(input, list):
|
||||
for item in input:
|
||||
if isinstance(item, dict):
|
||||
filter_value_from_dict(cast(dict, item), "cache_control")
|
||||
|
||||
if tools is not None:
|
||||
for tool in tools:
|
||||
if isinstance(tool, dict):
|
||||
filter_value_from_dict(cast(dict, tool), "cache_control")
|
||||
|
||||
return input, tools
|
||||
|
||||
def _validate_input_param(
|
||||
self, input: Union[str, ResponseInputParam]
|
||||
) -> Union[str, ResponseInputParam]:
|
||||
|
|
@ -604,6 +648,12 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
url = str(parsed_url.copy_with(path=compact_path))
|
||||
|
||||
input = self._validate_input_param(input)
|
||||
tools = response_api_optional_request_params.get("tools")
|
||||
input, tools = self.remove_cache_control_flag_from_input_and_tools(
|
||||
model=model, input=input, tools=tools
|
||||
)
|
||||
if tools is not None:
|
||||
response_api_optional_request_params["tools"] = tools
|
||||
data = dict(
|
||||
ResponsesAPIRequestParams(
|
||||
model=model, input=input, **response_api_optional_request_params
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
1
litellm/llms/reducto/__init__.py
Normal file
1
litellm/llms/reducto/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
|
||||
159
litellm/llms/reducto/common.py
Normal file
159
litellm/llms/reducto/common.py
Normal 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
|
||||
1
litellm/llms/reducto/ocr/__init__.py
Normal file
1
litellm/llms/reducto/ocr/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
|
||||
241
litellm/llms/reducto/ocr/transformation.py
Normal file
241
litellm/llms/reducto/ocr/transformation.py
Normal 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,
|
||||
)
|
||||
|
|
@ -578,7 +578,7 @@ class SagemakerLLM(BaseAWSLLM):
|
|||
logger_fn=None,
|
||||
):
|
||||
"""
|
||||
Supports both Huggingface Jumpstart embeddings and Voyage models
|
||||
Supports Hugging Face (TGI), Voyage, and Cohere embedding endpoints
|
||||
"""
|
||||
### BOTO3 INIT
|
||||
import boto3
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
141
litellm/llms/sagemaker/embedding/cohere_transformation.py
Normal file
141
litellm/llms/sagemaker/embedding/cohere_transformation.py
Normal file
|
|
@ -0,0 +1,141 @@
|
|||
"""
|
||||
Translate from OpenAI's `/v1/embeddings` to Sagemaker's `/invoke`
|
||||
|
||||
In the native Cohere embed format for self-hosted Cohere endpoints
|
||||
(AWS Marketplace / JumpStart). Cohere containers expect
|
||||
`{"texts": [...], "input_type": "..."}` and reject the HuggingFace TGI shape
|
||||
`{"inputs": [...]}` with `422 EmbedReqV2.inputs is of type string but should
|
||||
be of type Object`.
|
||||
|
||||
Reference: https://docs.cohere.com/v2/reference/embed
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, List, Optional, Union, cast
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.llms.openai import AllEmbeddingInputValues
|
||||
|
||||
from httpx._models import Headers, Response
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
|
||||
from litellm.llms.bedrock.embed.cohere_transformation import (
|
||||
BedrockCohereEmbeddingConfig,
|
||||
)
|
||||
from litellm.llms.cohere.embed.v1_transformation import CohereEmbeddingConfig
|
||||
from litellm.types.utils import EmbeddingResponse
|
||||
|
||||
from ..common_utils import SagemakerError
|
||||
|
||||
|
||||
class SagemakerCohereEmbeddingConfig(BaseEmbeddingConfig):
|
||||
"""
|
||||
SageMaker invoke payload for self-hosted Cohere embed models.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
return ["encoding_format", "dimensions", "input_type"]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
optional_params = BedrockCohereEmbeddingConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
if "input_type" in non_default_params:
|
||||
optional_params["input_type"] = non_default_params["input_type"]
|
||||
return optional_params
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, Headers]
|
||||
) -> BaseLLMException:
|
||||
return SagemakerError(
|
||||
message=error_message, status_code=status_code, headers=headers
|
||||
)
|
||||
|
||||
def transform_embedding_request(
|
||||
self,
|
||||
model: str,
|
||||
input: "AllEmbeddingInputValues",
|
||||
optional_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform embedding request for Cohere models on SageMaker
|
||||
"""
|
||||
if isinstance(input, str):
|
||||
input_list: List[str] = [input]
|
||||
elif isinstance(input, list):
|
||||
if input and (isinstance(input[0], list) or isinstance(input[0], int)):
|
||||
raise ValueError("Input must be a list of strings")
|
||||
input_list = cast(List[str], input)
|
||||
else:
|
||||
input_list = [str(input)]
|
||||
|
||||
return dict(
|
||||
BedrockCohereEmbeddingConfig()._transform_request(
|
||||
model=model,
|
||||
input=input_list,
|
||||
inference_params=optional_params,
|
||||
)
|
||||
)
|
||||
|
||||
def transform_embedding_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: Response,
|
||||
model_response: "EmbeddingResponse",
|
||||
logging_obj: Any,
|
||||
api_key: Optional[str] = None,
|
||||
request_data: dict = {},
|
||||
optional_params: dict = {},
|
||||
litellm_params: dict = {},
|
||||
) -> "EmbeddingResponse":
|
||||
"""
|
||||
Transform embedding response for Cohere models on SageMaker.
|
||||
|
||||
Uses `CohereEmbeddingConfig._populate_embedding_response` (not
|
||||
`_transform_response`) so we do not log `post_call` a second time
|
||||
— the SageMaker embedding handler already logs `post_call` before
|
||||
invoking this transform.
|
||||
"""
|
||||
input_value = (
|
||||
logging_obj.model_call_details.get("input")
|
||||
or request_data.get("texts")
|
||||
or request_data.get("images")
|
||||
or []
|
||||
)
|
||||
if isinstance(input_value, str):
|
||||
input_value = [input_value]
|
||||
|
||||
return CohereEmbeddingConfig()._populate_embedding_response(
|
||||
response_json=raw_response.json(),
|
||||
model_response=model_response,
|
||||
model=model,
|
||||
encoding=litellm.encoding,
|
||||
input=input_value,
|
||||
)
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
messages: List[Any],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""
|
||||
Validate environment for SageMaker Cohere embeddings
|
||||
"""
|
||||
return {"Content-Type": "application/json"}
|
||||
|
|
@ -1,7 +1,7 @@
|
|||
"""
|
||||
Translate from OpenAI's `/v1/embeddings` to Sagemaker's `/invoke`
|
||||
|
||||
In the Huggingface TGI format.
|
||||
In the Huggingface TGI format.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, List, Optional, Union
|
||||
|
|
@ -11,12 +11,13 @@ if TYPE_CHECKING:
|
|||
|
||||
from httpx._models import Headers, Response
|
||||
|
||||
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.types.utils import Usage, EmbeddingResponse
|
||||
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
|
||||
from litellm.llms.voyage.embedding.transformation import VoyageEmbeddingConfig
|
||||
from litellm.types.utils import EmbeddingResponse, Usage
|
||||
|
||||
from ..common_utils import SagemakerError
|
||||
from .cohere_transformation import SagemakerCohereEmbeddingConfig
|
||||
|
||||
|
||||
class SagemakerEmbeddingConfig(BaseEmbeddingConfig):
|
||||
|
|
@ -38,17 +39,20 @@ class SagemakerEmbeddingConfig(BaseEmbeddingConfig):
|
|||
Returns:
|
||||
Appropriate embedding config instance
|
||||
"""
|
||||
if "voyage" in model.lower():
|
||||
model_lower = model.lower()
|
||||
if "voyage" in model_lower:
|
||||
return VoyageEmbeddingConfig()
|
||||
else:
|
||||
return cls()
|
||||
if "cohere" in model_lower:
|
||||
return SagemakerCohereEmbeddingConfig()
|
||||
return cls()
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
# Check if this is an embedding model
|
||||
if "voyage" in model.lower():
|
||||
model_lower = model.lower()
|
||||
if "voyage" in model_lower:
|
||||
return VoyageEmbeddingConfig().get_supported_openai_params(model)
|
||||
else:
|
||||
return []
|
||||
if "cohere" in model_lower:
|
||||
return SagemakerCohereEmbeddingConfig().get_supported_openai_params(model)
|
||||
return []
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -41,7 +41,7 @@ class ContextCachingEndpoints(VertexBase):
|
|||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
super().__init__()
|
||||
|
||||
def _get_token_and_url_context_caching(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -1073,16 +1073,14 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915
|
|||
contents.append(ContentType(role="user", parts=tool_call_responses))
|
||||
|
||||
if len(contents) == 0:
|
||||
verbose_logger.warning(
|
||||
"""
|
||||
verbose_logger.warning("""
|
||||
No contents in messages. Contents are required. See
|
||||
https://cloud.google.com/vertex-ai/docs/reference/rest/v1/projects.locations.publishers.models/generateContent#request-body.
|
||||
If the original request did not comply to OpenAI API requirements it should have failed by now,
|
||||
but LiteLLM does not check for missing messages.
|
||||
Setting an empty content to prevent an 400 error.
|
||||
Relevant Issue - https://github.com/BerriAI/litellm/issues/9733
|
||||
"""
|
||||
)
|
||||
""")
|
||||
contents.append(ContentType(role="user", parts=[PartType(text=" ")]))
|
||||
return contents
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue