mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge branch 'litellm_internal_staging' into litellm_jwt_mapping_virtualkeys
This commit is contained in:
commit
17cf96a132
190 changed files with 10854 additions and 556 deletions
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 \
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -726,9 +726,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,
|
||||
|
|
@ -1617,6 +1665,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))
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
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="...")
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1409,6 +1409,8 @@ class BaseLLMHTTPHandler:
|
|||
document=document,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
# All providers return OCRRequestData
|
||||
|
|
@ -1477,6 +1479,8 @@ class BaseLLMHTTPHandler:
|
|||
document=document,
|
||||
optional_params=optional_params,
|
||||
headers=headers,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
# All providers return OCRRequestData
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -31,7 +31,7 @@ from ..vertex_llm_base import VertexBase
|
|||
|
||||
class VertexAIGemmaModels(VertexBase):
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
super().__init__()
|
||||
|
||||
def completion(
|
||||
self,
|
||||
|
|
@ -62,9 +62,6 @@ class VertexAIGemmaModels(VertexBase):
|
|||
try:
|
||||
import vertexai
|
||||
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
VertexLLM,
|
||||
)
|
||||
from litellm.llms.vertex_ai.vertex_gemma_models.transformation import (
|
||||
VertexGemmaConfig,
|
||||
)
|
||||
|
|
@ -83,9 +80,8 @@ class VertexAIGemmaModels(VertexBase):
|
|||
)
|
||||
try:
|
||||
model = get_vertex_base_model_name(model=model)
|
||||
vertex_httpx_logic = VertexLLM()
|
||||
|
||||
access_token, project_id = vertex_httpx_logic._ensure_access_token(
|
||||
access_token, project_id = self._ensure_access_token(
|
||||
credentials=vertex_credentials,
|
||||
project_id=vertex_project,
|
||||
custom_llm_provider="vertex_ai",
|
||||
|
|
|
|||
|
|
@ -91,6 +91,10 @@ class VertexGemmaConfig(OpenAIGPTConfig):
|
|||
"stream", None
|
||||
) # Streaming not supported, will be faked client-side
|
||||
openai_request.pop("stream_options", None) # Stream options not supported
|
||||
# Vertex Gemma's chatCompletions wrapper does not understand
|
||||
# `context_management` (an Anthropic/Responses API concept). Strip it
|
||||
# so the upstream endpoint does not 400 on the unknown field.
|
||||
openai_request.pop("context_management", None)
|
||||
|
||||
# Wrap in Vertex Gemma format
|
||||
return {
|
||||
|
|
|
|||
|
|
@ -4,8 +4,10 @@ Base Vertex, Google AI Studio LLM Class
|
|||
Handles Authentication and generating request urls for Vertex AI and Google AI Studio
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import threading
|
||||
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple
|
||||
|
||||
import litellm
|
||||
|
|
@ -30,6 +32,7 @@ GOOGLE_IMPORT_ERROR_MESSAGE = (
|
|||
|
||||
if TYPE_CHECKING:
|
||||
from google.auth.credentials import Credentials as GoogleCredentialsObject
|
||||
from google.auth.credentials import TokenState
|
||||
else:
|
||||
GoogleCredentialsObject = Any
|
||||
|
||||
|
|
@ -42,10 +45,28 @@ class VertexBase:
|
|||
self._credentials: Optional[GoogleCredentialsObject] = None
|
||||
self._credentials_project_mapping: Dict[
|
||||
Tuple[Optional[VERTEX_CREDENTIALS_TYPES], Optional[str]],
|
||||
Tuple[GoogleCredentialsObject, str],
|
||||
Tuple[GoogleCredentialsObject, Optional[str]],
|
||||
] = {}
|
||||
self.project_id: Optional[str] = None
|
||||
self.async_handler: Optional[AsyncHTTPHandler] = None
|
||||
# Per-credential-key asyncio.Lock for single-flight async refresh.
|
||||
# Prevents thundering herd when token expires under high concurrency.
|
||||
# Uses a regular dict (not WeakValueDictionary) so the lock identity is
|
||||
# stable across concurrent callers — a weak reference can be GC'd
|
||||
# between two coroutines arriving at the lock, breaking single-flight.
|
||||
# An explicit refcount tracks the number of coroutines currently using
|
||||
# each lock; the entry is pruned when the count reaches zero, so the
|
||||
# dict stays bounded even in long-running high-cardinality deployments
|
||||
# without depending on any private asyncio internals.
|
||||
self._async_refresh_locks: Dict[tuple, asyncio.Lock] = {}
|
||||
self._async_refresh_lock_refcounts: Dict[tuple, int] = {}
|
||||
# Tracks in-flight background refresh tasks to avoid duplicate refreshes.
|
||||
self._background_refresh_tasks: Dict[tuple, asyncio.Task] = {}
|
||||
# Protects the sync get_access_token refresh path.
|
||||
# Use RLock so that the reauthentication retry path (which calls
|
||||
# back into get_access_token while still holding the lock) can
|
||||
# re-acquire it without deadlocking the current thread.
|
||||
self._sync_refresh_lock = threading.RLock()
|
||||
|
||||
def get_vertex_region(self, vertex_region: Optional[str], model: str) -> str:
|
||||
import litellm
|
||||
|
|
@ -77,7 +98,9 @@ class VertexBase:
|
|||
return vertex_region or "us-central1"
|
||||
|
||||
def load_auth(
|
||||
self, credentials: Optional[VERTEX_CREDENTIALS_TYPES], project_id: Optional[str]
|
||||
self,
|
||||
credentials: Optional[VERTEX_CREDENTIALS_TYPES],
|
||||
project_id: Optional[str],
|
||||
) -> Tuple[Any, str]:
|
||||
if credentials is not None:
|
||||
if isinstance(credentials, str):
|
||||
|
|
@ -343,7 +366,241 @@ class VertexBase:
|
|||
except ImportError:
|
||||
raise ImportError(GOOGLE_IMPORT_ERROR_MESSAGE)
|
||||
|
||||
credentials.refresh(Request())
|
||||
# Serialize all refreshes on this VertexBase across threads.
|
||||
# ``credentials.refresh()`` is not safe to call concurrently on the
|
||||
# same credentials object, and this method is invoked from three
|
||||
# places that can run on different threads:
|
||||
# - sync ``get_access_token`` (already holds ``_sync_refresh_lock``)
|
||||
# - the async slow path (via ``asyncify`` in a worker thread)
|
||||
# - the background proactive refresh task (via ``asyncify``)
|
||||
# ``_sync_refresh_lock`` is an ``RLock`` so reentrant acquisition
|
||||
# from the sync path is safe.
|
||||
with self._sync_refresh_lock:
|
||||
credentials.refresh(Request())
|
||||
|
||||
def _acquire_async_refresh_lock(self, credential_cache_key: tuple) -> asyncio.Lock:
|
||||
"""Increment the refcount and return the lock for ``credential_cache_key``.
|
||||
|
||||
Every call must be paired with ``_release_async_refresh_lock`` once the
|
||||
caller is done with the lock so the entry can be pruned when no other
|
||||
coroutine is holding or waiting on it.
|
||||
"""
|
||||
lock = self._async_refresh_locks.setdefault(
|
||||
credential_cache_key, asyncio.Lock()
|
||||
)
|
||||
self._async_refresh_lock_refcounts[credential_cache_key] = (
|
||||
self._async_refresh_lock_refcounts.get(credential_cache_key, 0) + 1
|
||||
)
|
||||
return lock
|
||||
|
||||
def _release_async_refresh_lock(
|
||||
self, credential_cache_key: tuple, lock: asyncio.Lock
|
||||
) -> None:
|
||||
"""Decrement the refcount and drop the lock entry when it reaches zero.
|
||||
|
||||
Must be called only after the caller has released ``lock`` (i.e. once
|
||||
the surrounding ``async with`` has exited). asyncio is cooperative, so
|
||||
the decrement-then-pop sequence below runs atomically with respect to
|
||||
other coroutines.
|
||||
"""
|
||||
remaining = self._async_refresh_lock_refcounts.get(credential_cache_key, 0) - 1
|
||||
if remaining > 0:
|
||||
self._async_refresh_lock_refcounts[credential_cache_key] = remaining
|
||||
return
|
||||
self._async_refresh_lock_refcounts.pop(credential_cache_key, None)
|
||||
if self._async_refresh_locks.get(credential_cache_key) is lock:
|
||||
self._async_refresh_locks.pop(credential_cache_key, None)
|
||||
|
||||
def _try_get_cached_token(
|
||||
self,
|
||||
credential_cache_key: tuple,
|
||||
project_id: Optional[str],
|
||||
) -> Optional[Tuple[str, str]]:
|
||||
"""
|
||||
Look up cached credentials and return (token, project_id) if the token
|
||||
is FRESH. Returns None if not cached or not fresh.
|
||||
"""
|
||||
from google.auth.credentials import TokenState
|
||||
|
||||
creds, cached_project_id = self._unpack_cached_credentials(credential_cache_key)
|
||||
if (
|
||||
creds is not None
|
||||
and self._get_token_state(creds) == TokenState.FRESH
|
||||
and creds.token is not None
|
||||
and isinstance(creds.token, str)
|
||||
):
|
||||
resolved_project = project_id or cached_project_id
|
||||
if resolved_project:
|
||||
return creds.token, resolved_project
|
||||
return None
|
||||
|
||||
def _try_get_usable_cached_token(
|
||||
self,
|
||||
credential_cache_key: tuple,
|
||||
project_id: Optional[str],
|
||||
) -> Optional[Tuple[str, str, "TokenState", Any, Optional[str]]]:
|
||||
"""
|
||||
Look up cached credentials and return usable token info for FRESH or
|
||||
STALE tokens (both are still valid for outbound requests). STALE
|
||||
tokens are returned along with their state and the underlying
|
||||
credentials object so the caller can schedule a background refresh
|
||||
without holding the per-key async lock.
|
||||
"""
|
||||
from google.auth.credentials import TokenState
|
||||
|
||||
creds, cached_project_id = self._unpack_cached_credentials(credential_cache_key)
|
||||
if creds is None:
|
||||
return None
|
||||
token_state = self._get_token_state(creds)
|
||||
if token_state not in (TokenState.FRESH, TokenState.STALE):
|
||||
return None
|
||||
if creds.token is None or not isinstance(creds.token, str):
|
||||
return None
|
||||
resolved_project = project_id or cached_project_id
|
||||
if not resolved_project:
|
||||
return None
|
||||
return creds.token, resolved_project, token_state, creds, cached_project_id
|
||||
|
||||
def _unpack_cached_credentials(
|
||||
self, credential_cache_key: tuple
|
||||
) -> Tuple[Any, Optional[str]]:
|
||||
"""
|
||||
Return (credentials, project_id) from the cache, or (None, None) if
|
||||
not cached. Handles both tuple and legacy cache formats.
|
||||
"""
|
||||
if credential_cache_key not in self._credentials_project_mapping:
|
||||
return None, None
|
||||
cached_entry = self._credentials_project_mapping[credential_cache_key]
|
||||
if isinstance(cached_entry, tuple):
|
||||
return cached_entry
|
||||
return cached_entry, cached_entry.quota_project_id or getattr(
|
||||
cached_entry, "project_id", None
|
||||
)
|
||||
|
||||
def _get_token_state(self, credentials: Any) -> "TokenState":
|
||||
"""
|
||||
Return the token state using google-auth's TokenState enum.
|
||||
|
||||
Falls back to expired/valid checks if token_state is unavailable
|
||||
(e.g. older google-auth versions or mock objects in tests).
|
||||
"""
|
||||
from google.auth.credentials import TokenState as _TokenState
|
||||
|
||||
token_state = getattr(credentials, "token_state", None)
|
||||
if isinstance(token_state, _TokenState):
|
||||
return token_state
|
||||
# Fallback for credentials without a real token_state (e.g. mocks)
|
||||
if getattr(credentials, "expired", True):
|
||||
return _TokenState.INVALID
|
||||
if getattr(credentials, "valid", False):
|
||||
return _TokenState.FRESH
|
||||
return _TokenState.INVALID
|
||||
|
||||
async def _load_and_cache_credentials(
|
||||
self,
|
||||
credentials: Optional[VERTEX_CREDENTIALS_TYPES],
|
||||
project_id: Optional[str],
|
||||
credential_cache_key: tuple,
|
||||
) -> Tuple[Any, Optional[str]]:
|
||||
"""Load credentials via load_auth (in thread) and cache the result."""
|
||||
try:
|
||||
_credentials, credential_project_id = await asyncify(self.load_auth)(
|
||||
credentials=credentials,
|
||||
project_id=project_id,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.exception("Failed to load vertex credentials: %s", str(e))
|
||||
raise
|
||||
if _credentials is None:
|
||||
raise ValueError("Could not resolve credentials")
|
||||
self._credentials_project_mapping[credential_cache_key] = (
|
||||
_credentials,
|
||||
credential_project_id,
|
||||
)
|
||||
return _credentials, credential_project_id
|
||||
|
||||
async def _background_refresh_credentials(
|
||||
self,
|
||||
credentials: Any,
|
||||
credential_cache_key: tuple,
|
||||
credential_project_id: Optional[str],
|
||||
) -> None:
|
||||
"""
|
||||
Refresh credentials in the background without blocking the calling request.
|
||||
|
||||
Called when the token is still valid but nearing expiry (proactive refresh).
|
||||
Errors are logged but not raised — the current token is still usable.
|
||||
"""
|
||||
try:
|
||||
verbose_logger.debug("Background proactive credential refresh")
|
||||
await asyncify(self.refresh_auth)(credentials)
|
||||
# Only update the cache if it still points at the credentials
|
||||
# object we just refreshed. The per-key async lock is not held
|
||||
# here, so a concurrent INVALID path may have already replaced
|
||||
# this entry (e.g. via _handle_reauthentication_async, which
|
||||
# creates a fresh credentials object). In that case our write
|
||||
# would clobber the newer entry with a stale reference.
|
||||
cached_creds, _ = self._unpack_cached_credentials(credential_cache_key)
|
||||
if cached_creds is credentials:
|
||||
self._credentials_project_mapping[credential_cache_key] = (
|
||||
credentials,
|
||||
credential_project_id,
|
||||
)
|
||||
except Exception:
|
||||
verbose_logger.debug(
|
||||
"Background credential refresh failed, will retry on next request",
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
async def _await_in_flight_background_refresh(
|
||||
self, credential_cache_key: tuple
|
||||
) -> None:
|
||||
"""Wait for an in-flight background refresh to finish, if any.
|
||||
|
||||
google-auth's ``Credentials.refresh()`` is not safe to invoke
|
||||
concurrently on the same credentials object. Coroutines that need a
|
||||
blocking refresh must first drain any background refresh that was
|
||||
scheduled while a previous STALE token was being served.
|
||||
"""
|
||||
existing_task = self._background_refresh_tasks.get(credential_cache_key)
|
||||
if existing_task is None or existing_task.done():
|
||||
return
|
||||
try:
|
||||
await existing_task
|
||||
except Exception:
|
||||
# Background refresh failures are already logged inside
|
||||
# _background_refresh_credentials; the caller will fall through
|
||||
# to its own blocking refresh.
|
||||
pass
|
||||
|
||||
def _schedule_background_refresh(
|
||||
self,
|
||||
credentials: Any,
|
||||
credential_cache_key: tuple,
|
||||
credential_project_id: Optional[str],
|
||||
) -> None:
|
||||
"""Kick off a single background refresh for ``credential_cache_key``.
|
||||
|
||||
Skips scheduling if a refresh is already in flight. The done-callback
|
||||
guards against removing a newer task that has replaced this one in the
|
||||
tracking dict (done_callbacks are scheduled via ``call_soon``).
|
||||
"""
|
||||
existing = self._background_refresh_tasks.get(credential_cache_key)
|
||||
if existing is not None and not existing.done():
|
||||
return
|
||||
self._background_refresh_tasks.pop(credential_cache_key, None)
|
||||
task = asyncio.create_task(
|
||||
self._background_refresh_credentials(
|
||||
credentials, credential_cache_key, credential_project_id
|
||||
)
|
||||
)
|
||||
|
||||
def _drop_background_refresh_task(_fut: asyncio.Future[Any]) -> None:
|
||||
if self._background_refresh_tasks.get(credential_cache_key) is _fut:
|
||||
self._background_refresh_tasks.pop(credential_cache_key, None)
|
||||
|
||||
task.add_done_callback(_drop_background_refresh_task)
|
||||
self._background_refresh_tasks[credential_cache_key] = task
|
||||
|
||||
def _ensure_access_token(
|
||||
self,
|
||||
|
|
@ -563,6 +820,65 @@ class VertexBase:
|
|||
# Re-raise the original error for better context
|
||||
raise error
|
||||
|
||||
async def _handle_reauthentication_async(
|
||||
self,
|
||||
credentials: Optional[VERTEX_CREDENTIALS_TYPES],
|
||||
project_id: Optional[str],
|
||||
credential_cache_key: Tuple,
|
||||
error: Exception,
|
||||
) -> Tuple[str, str]:
|
||||
"""
|
||||
Async reauthentication retry that stays within the per-key async lock.
|
||||
"""
|
||||
verbose_logger.debug(
|
||||
f"Handling async reauthentication for project_id: {project_id}. "
|
||||
f"Clearing cache and retrying once."
|
||||
)
|
||||
|
||||
self._credentials_project_mapping.pop(credential_cache_key, None)
|
||||
|
||||
try:
|
||||
_credentials, credential_project_id = (
|
||||
await self._load_and_cache_credentials(
|
||||
credentials=credentials,
|
||||
project_id=project_id,
|
||||
credential_cache_key=credential_cache_key,
|
||||
)
|
||||
)
|
||||
if project_id is None and isinstance(credential_project_id, str):
|
||||
project_id = credential_project_id
|
||||
cache_credentials = (
|
||||
json.dumps(credentials)
|
||||
if isinstance(credentials, dict)
|
||||
else credentials
|
||||
)
|
||||
resolved_cache_key = (cache_credentials, project_id)
|
||||
# Always overwrite — any pre-existing entry at the resolved key
|
||||
# references the OLD credentials object we just replaced, and
|
||||
# leaving it would force the next request to do a redundant
|
||||
# refresh/reauth before realizing the cached creds are stale.
|
||||
self._credentials_project_mapping[resolved_cache_key] = (
|
||||
_credentials,
|
||||
credential_project_id,
|
||||
)
|
||||
|
||||
if _credentials.token is None or not isinstance(_credentials.token, str):
|
||||
raise ValueError(
|
||||
"Could not resolve credentials token. Got None or non-string token (type={})".format(
|
||||
type(_credentials.token).__name__
|
||||
)
|
||||
)
|
||||
if project_id is None:
|
||||
raise ValueError("Could not resolve project_id")
|
||||
|
||||
return _credentials.token, project_id
|
||||
except Exception as retry_error:
|
||||
verbose_logger.error(
|
||||
f"Async reauthentication retry failed for project_id: {project_id}. "
|
||||
f"Original error: {str(error)}. Retry error: {str(retry_error)}"
|
||||
)
|
||||
raise error
|
||||
|
||||
def get_access_token(
|
||||
self,
|
||||
credentials: Optional[VERTEX_CREDENTIALS_TYPES],
|
||||
|
|
@ -646,7 +962,7 @@ class VertexBase:
|
|||
)
|
||||
|
||||
## VALIDATE CREDENTIALS
|
||||
verbose_logger.debug(f"Validating credentials for project_id: {project_id}")
|
||||
verbose_logger.debug("Validating credentials")
|
||||
if (
|
||||
project_id is None
|
||||
and credential_project_id is not None
|
||||
|
|
@ -666,26 +982,27 @@ class VertexBase:
|
|||
raise ValueError("Credentials are None after loading")
|
||||
|
||||
if _credentials.expired:
|
||||
try:
|
||||
verbose_logger.debug(
|
||||
f"Credentials expired, refreshing for project_id: {project_id}"
|
||||
)
|
||||
self.refresh_auth(_credentials)
|
||||
self._credentials_project_mapping[credential_cache_key] = (
|
||||
_credentials,
|
||||
credential_project_id,
|
||||
)
|
||||
except Exception as e:
|
||||
# if refresh fails, it's possible the user has re-authenticated via `gcloud auth application-default login`
|
||||
# in this case, we should try to reload the credentials by clearing the cache and retrying
|
||||
if "Reauthentication is needed" in str(e) and not _retry_reauth:
|
||||
return self._handle_reauthentication(
|
||||
credentials=credentials,
|
||||
project_id=project_id,
|
||||
credential_cache_key=credential_cache_key,
|
||||
error=e,
|
||||
)
|
||||
raise e
|
||||
with self._sync_refresh_lock:
|
||||
# Double-check after acquiring lock
|
||||
if _credentials.expired:
|
||||
try:
|
||||
verbose_logger.debug("Credentials expired, refreshing")
|
||||
self.refresh_auth(_credentials)
|
||||
self._credentials_project_mapping[credential_cache_key] = (
|
||||
_credentials,
|
||||
credential_project_id,
|
||||
)
|
||||
except Exception as e:
|
||||
# if refresh fails, it's possible the user has re-authenticated via `gcloud auth application-default login`
|
||||
# in this case, we should try to reload the credentials by clearing the cache and retrying
|
||||
if "Reauthentication is needed" in str(e) and not _retry_reauth:
|
||||
return self._handle_reauthentication(
|
||||
credentials=credentials,
|
||||
project_id=project_id,
|
||||
credential_cache_key=credential_cache_key,
|
||||
error=e,
|
||||
)
|
||||
raise e
|
||||
|
||||
## VALIDATION STEP
|
||||
if _credentials.token is None or not isinstance(_credentials.token, str):
|
||||
|
|
@ -700,6 +1017,149 @@ class VertexBase:
|
|||
|
||||
return _credentials.token, project_id
|
||||
|
||||
async def get_access_token_async(
|
||||
self,
|
||||
credentials: Optional[VERTEX_CREDENTIALS_TYPES],
|
||||
project_id: Optional[str],
|
||||
) -> Tuple[str, str]:
|
||||
"""
|
||||
Async version of get_access_token with single-flight refresh coordination.
|
||||
|
||||
Prevents thundering herd: when credentials expire under high concurrency,
|
||||
only one coroutine refreshes while others wait on the lock. Uses native
|
||||
async refresh for service_account and authorized_user credentials.
|
||||
"""
|
||||
from google.auth.credentials import TokenState
|
||||
|
||||
cache_credentials = (
|
||||
json.dumps(credentials) if isinstance(credentials, dict) else credentials
|
||||
)
|
||||
credential_cache_key = (cache_credentials, project_id)
|
||||
|
||||
# === FAST PATH (no lock) ===
|
||||
# If credentials are FRESH or STALE, return immediately without
|
||||
# touching the per-key async lock. STALE tokens are still usable;
|
||||
# we kick off a deduplicated background refresh so subsequent
|
||||
# requests get a fresh token, but we must not serialize concurrent
|
||||
# callers on the lock just to schedule that refresh.
|
||||
usable = self._try_get_usable_cached_token(credential_cache_key, project_id)
|
||||
if usable is not None:
|
||||
cached_token, resolved_project, token_state, creds, cached_project_id = (
|
||||
usable
|
||||
)
|
||||
if token_state == TokenState.STALE:
|
||||
self._schedule_background_refresh(
|
||||
creds, credential_cache_key, cached_project_id
|
||||
)
|
||||
return cached_token, resolved_project
|
||||
|
||||
# === SLOW PATH (per-key lock) ===
|
||||
lock = self._acquire_async_refresh_lock(credential_cache_key)
|
||||
try:
|
||||
async with lock:
|
||||
# Double-check after acquiring lock — another coroutine may have refreshed.
|
||||
cached = self._try_get_cached_token(credential_cache_key, project_id)
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
_credentials, credential_project_id = self._unpack_cached_credentials(
|
||||
credential_cache_key
|
||||
)
|
||||
|
||||
# Load credentials if not cached
|
||||
if _credentials is None:
|
||||
_credentials, credential_project_id = (
|
||||
await self._load_and_cache_credentials(
|
||||
credentials, project_id, credential_cache_key
|
||||
)
|
||||
)
|
||||
|
||||
# Resolve project_id from credentials if not provided
|
||||
if project_id is None and isinstance(credential_project_id, str):
|
||||
project_id = credential_project_id
|
||||
resolved_cache_key = (cache_credentials, project_id)
|
||||
# Always overwrite — a pre-existing entry at the resolved
|
||||
# key may reference stale credentials (e.g. from before a
|
||||
# reauth that only repopulated the unresolved key), which
|
||||
# would force the next request through an unnecessary
|
||||
# refresh/reauth cycle.
|
||||
self._credentials_project_mapping[resolved_cache_key] = (
|
||||
_credentials,
|
||||
credential_project_id,
|
||||
)
|
||||
|
||||
# Use google-auth's token_state to decide refresh strategy:
|
||||
# - STALE: token is usable but within REFRESH_THRESHOLD (3:45) of
|
||||
# expiry — return it immediately and refresh in the background.
|
||||
# - INVALID: token is expired or missing — must block on refresh.
|
||||
token_state = self._get_token_state(_credentials)
|
||||
|
||||
if token_state == TokenState.STALE:
|
||||
if project_id is None:
|
||||
raise ValueError("Could not resolve project_id")
|
||||
current_token = _credentials.token
|
||||
if current_token is None or not isinstance(current_token, str):
|
||||
# Token is malformed despite STALE state — block on a full
|
||||
# refresh using the same path as INVALID credentials.
|
||||
token_state = TokenState.INVALID
|
||||
else:
|
||||
self._schedule_background_refresh(
|
||||
_credentials,
|
||||
credential_cache_key,
|
||||
credential_project_id,
|
||||
)
|
||||
return current_token, project_id
|
||||
|
||||
if token_state == TokenState.INVALID:
|
||||
# Drain any in-flight background refresh before invoking
|
||||
# refresh_auth ourselves; google-auth's
|
||||
# Credentials.refresh() is not safe to call concurrently
|
||||
# on the same credentials object, and the background task
|
||||
# runs outside this lock.
|
||||
await self._await_in_flight_background_refresh(credential_cache_key)
|
||||
cached = self._try_get_cached_token(
|
||||
credential_cache_key, project_id
|
||||
)
|
||||
if cached is not None:
|
||||
return cached
|
||||
|
||||
# Token is expired or missing — must block until refresh completes.
|
||||
try:
|
||||
verbose_logger.debug("Credentials expired, refreshing")
|
||||
await asyncify(self.refresh_auth)(_credentials)
|
||||
self._credentials_project_mapping[credential_cache_key] = (
|
||||
_credentials,
|
||||
credential_project_id,
|
||||
)
|
||||
except Exception as e:
|
||||
if "Reauthentication is needed" in str(e):
|
||||
verbose_logger.debug(
|
||||
"Reauthentication needed, clearing cache and retrying"
|
||||
)
|
||||
return await self._handle_reauthentication_async(
|
||||
credentials=credentials,
|
||||
project_id=project_id,
|
||||
credential_cache_key=credential_cache_key,
|
||||
error=e,
|
||||
)
|
||||
raise
|
||||
|
||||
# Final validation
|
||||
if _credentials.token is None or not isinstance(
|
||||
_credentials.token, str
|
||||
):
|
||||
raise ValueError(
|
||||
"Could not resolve credentials token. Got None or non-string token (type={})".format(
|
||||
type(_credentials.token).__name__
|
||||
)
|
||||
)
|
||||
if project_id is None:
|
||||
raise ValueError("Could not resolve project_id")
|
||||
|
||||
return _credentials.token, project_id
|
||||
finally:
|
||||
self._release_async_refresh_lock(credential_cache_key, lock)
|
||||
|
||||
async def _ensure_access_token_async(
|
||||
self,
|
||||
credentials: Optional[VERTEX_CREDENTIALS_TYPES],
|
||||
|
|
@ -714,13 +1174,10 @@ class VertexBase:
|
|||
if custom_llm_provider == "gemini":
|
||||
return "", ""
|
||||
else:
|
||||
try:
|
||||
return await asyncify(self.get_access_token)(
|
||||
credentials=credentials,
|
||||
project_id=project_id,
|
||||
)
|
||||
except Exception as e:
|
||||
raise e
|
||||
return await self.get_access_token_async(
|
||||
credentials=credentials,
|
||||
project_id=project_id,
|
||||
)
|
||||
|
||||
def set_headers(
|
||||
self, auth_header: Optional[str], extra_headers: Optional[dict]
|
||||
|
|
|
|||
|
|
@ -57,7 +57,7 @@ def create_vertex_url(
|
|||
|
||||
class VertexAIModelGardenModels(VertexBase):
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
super().__init__()
|
||||
|
||||
def completion(
|
||||
self,
|
||||
|
|
@ -89,9 +89,6 @@ class VertexAIModelGardenModels(VertexBase):
|
|||
import vertexai
|
||||
|
||||
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,
|
||||
|
|
@ -107,9 +104,8 @@ class VertexAIModelGardenModels(VertexBase):
|
|||
)
|
||||
try:
|
||||
model = get_vertex_base_model_name(model=model)
|
||||
vertex_httpx_logic = VertexLLM()
|
||||
|
||||
access_token, project_id = vertex_httpx_logic._ensure_access_token(
|
||||
access_token, project_id = self._ensure_access_token(
|
||||
credentials=vertex_credentials,
|
||||
project_id=vertex_project,
|
||||
custom_llm_provider="vertex_ai",
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Translates from OpenAI's `/v1/chat/completions` to the VLLM sdk `llm.generate`.
|
||||
Translates from OpenAI's `/v1/chat/completions` to the VLLM sdk `llm.generate`.
|
||||
|
||||
NOT RECOMMENDED FOR PRODUCTION USE. Use `hosted_vllm/` instead.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
"""
|
||||
This module is used to transform the request and response for the Voyage contextualized embeddings API.
|
||||
This would be used for all the contextualized embeddings models in Voyage.
|
||||
This module is used to transform the request and response for the Voyage contextualized embeddings API.
|
||||
This would be used for all the contextualized embeddings models in Voyage.
|
||||
"""
|
||||
|
||||
from typing import List, Optional, Union
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import Any, AsyncIterator, Iterator, List, Optional, Tuple, Union
|
||||
from typing import Any, AsyncIterator, Dict, Iterator, List, Optional, Tuple, Union
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -26,6 +26,7 @@ from ...openai.chat.gpt_transformation import (
|
|||
|
||||
|
||||
class XAIChatConfig(OpenAIGPTConfig):
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "xai"
|
||||
|
|
@ -225,21 +226,57 @@ class XAIChatConfig(OpenAIGPTConfig):
|
|||
verbose_logger.debug(f"Error extracting X.AI web search usage: {e}")
|
||||
|
||||
self._fold_reasoning_tokens_into_completion(response)
|
||||
self._normalize_openai_compatible_usage_totals(getattr(response, "usage", None))
|
||||
return response
|
||||
|
||||
@staticmethod
|
||||
def _fold_reasoning_tokens_into_completion(model_response: ModelResponse) -> None:
|
||||
def _fold_reasoning_tokens_into_completion(
|
||||
target: Union[ModelResponse, Usage, Dict[str, Any], None],
|
||||
) -> None:
|
||||
"""Reconcile xAI Usage to the OpenAI invariant.
|
||||
|
||||
xAI accounts ``reasoning_tokens`` separately from
|
||||
``completion_tokens`` while still summing them into ``total_tokens``.
|
||||
OpenAI's contract (o1/o3) folds reasoning into ``completion_tokens``,
|
||||
so fold here to keep ``total = prompt + completion``. Idempotent.
|
||||
|
||||
Accepts a ``ModelResponse`` (non-streaming), a ``Usage`` object, or a
|
||||
raw usage ``dict`` (streaming chunk) so streaming and non-streaming
|
||||
paths stay in sync.
|
||||
"""
|
||||
usage = getattr(model_response, "usage", None)
|
||||
if target is None:
|
||||
return
|
||||
|
||||
if isinstance(target, ModelResponse):
|
||||
usage: Union[Usage, Dict[str, Any], None] = getattr(target, "usage", None)
|
||||
else:
|
||||
usage = target
|
||||
if usage is None:
|
||||
return
|
||||
|
||||
if isinstance(usage, dict):
|
||||
details = usage.get("completion_tokens_details") or {}
|
||||
if isinstance(details, dict):
|
||||
reasoning_tokens = int(details.get("reasoning_tokens") or 0)
|
||||
else:
|
||||
reasoning_tokens = int(getattr(details, "reasoning_tokens", 0) or 0)
|
||||
if reasoning_tokens <= 0:
|
||||
return
|
||||
|
||||
prompt_tokens = int(usage.get("prompt_tokens") or 0)
|
||||
completion_tokens = int(usage.get("completion_tokens") or 0)
|
||||
total_tokens = int(usage.get("total_tokens") or 0)
|
||||
|
||||
if total_tokens == prompt_tokens + completion_tokens:
|
||||
return
|
||||
|
||||
# Guard against double-counting if xAI changes accounting.
|
||||
if total_tokens != prompt_tokens + completion_tokens + reasoning_tokens:
|
||||
return
|
||||
|
||||
usage["completion_tokens"] = completion_tokens + reasoning_tokens
|
||||
return
|
||||
|
||||
details = getattr(usage, "completion_tokens_details", None)
|
||||
reasoning_tokens = (
|
||||
int(getattr(details, "reasoning_tokens", 0) or 0) if details else 0
|
||||
|
|
@ -284,6 +321,25 @@ class XAIChatConfig(OpenAIGPTConfig):
|
|||
setattr(usage, "num_sources_used", int(num_sources_used))
|
||||
verbose_logger.debug(f"X.AI web search sources used: {num_sources_used}")
|
||||
|
||||
@staticmethod
|
||||
def _normalize_openai_compatible_usage_totals(
|
||||
usage: Union[Usage, Dict[str, Any], None],
|
||||
) -> None:
|
||||
if usage is None:
|
||||
return
|
||||
if isinstance(usage, dict):
|
||||
prompt_tokens = int(usage.get("prompt_tokens") or 0)
|
||||
completion_tokens = int(usage.get("completion_tokens") or 0)
|
||||
expected_total = prompt_tokens + completion_tokens
|
||||
if int(usage.get("total_tokens") or 0) < expected_total:
|
||||
usage["total_tokens"] = expected_total
|
||||
return
|
||||
prompt_tokens = int(usage.prompt_tokens or 0)
|
||||
completion_tokens = int(usage.completion_tokens or 0)
|
||||
expected_total = prompt_tokens + completion_tokens
|
||||
if int(usage.total_tokens or 0) < expected_total:
|
||||
usage.total_tokens = expected_total
|
||||
|
||||
|
||||
class XAIChatCompletionStreamingHandler(OpenAIChatCompletionStreamingHandler):
|
||||
def chunk_parser(self, chunk: dict) -> ModelResponseStream:
|
||||
|
|
@ -304,4 +360,8 @@ class XAIChatCompletionStreamingHandler(OpenAIChatCompletionStreamingHandler):
|
|||
# Add a dummy choice with empty delta to ensure proper processing
|
||||
chunk["choices"] = [{"index": 0, "delta": {}, "finish_reason": None}]
|
||||
|
||||
if "usage" in chunk and chunk["usage"] is not None:
|
||||
XAIChatConfig._fold_reasoning_tokens_into_completion(chunk["usage"])
|
||||
XAIChatConfig._normalize_openai_compatible_usage_totals(chunk["usage"])
|
||||
|
||||
return super().chunk_parser(chunk)
|
||||
|
|
|
|||
|
|
@ -13982,6 +13982,21 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/glm-5p1": {
|
||||
"cache_read_input_token_cost": 2.6e-07,
|
||||
"input_cost_per_token": 1.4e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 202800,
|
||||
"max_output_tokens": 202800,
|
||||
"max_tokens": 202800,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 4.4e-06,
|
||||
"source": "https://fireworks.ai/models/fireworks/glm-5p1",
|
||||
"supports_function_calling": false,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": false,
|
||||
"supports_tool_choice": false
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/gpt-oss-120b": {
|
||||
"input_cost_per_token": 1.5e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
|
|
@ -14248,6 +14263,21 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"fireworks_ai/glm-5p1": {
|
||||
"cache_read_input_token_cost": 2.6e-07,
|
||||
"input_cost_per_token": 1.4e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 202800,
|
||||
"max_output_tokens": 202800,
|
||||
"max_tokens": 202800,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 4.4e-06,
|
||||
"source": "https://fireworks.ai/models/fireworks/glm-5p1",
|
||||
"supports_function_calling": false,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": false,
|
||||
"supports_tool_choice": false
|
||||
},
|
||||
"fireworks_ai/kimi-k2p5": {
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
"input_cost_per_token": 6e-07,
|
||||
|
|
@ -29122,6 +29152,24 @@
|
|||
"supports_tool_choice": true,
|
||||
"source": "https://aws.amazon.com/bedrock/pricing/"
|
||||
},
|
||||
"reducto/parse-legacy": {
|
||||
"litellm_provider": "reducto",
|
||||
"mode": "ocr",
|
||||
"ocr_cost_per_credit": 0.015,
|
||||
"source": "https://reducto.ai/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/ocr"
|
||||
]
|
||||
},
|
||||
"reducto/parse-v3": {
|
||||
"litellm_provider": "reducto",
|
||||
"mode": "ocr",
|
||||
"ocr_cost_per_credit": 0.015,
|
||||
"source": "https://reducto.ai/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/ocr"
|
||||
]
|
||||
},
|
||||
"recraft/recraftv2": {
|
||||
"litellm_provider": "recraft",
|
||||
"mode": "image_generation",
|
||||
|
|
|
|||
|
|
@ -305,7 +305,7 @@ def build_input_schema(operation: Dict[str, Any]) -> Dict[str, Any]:
|
|||
|
||||
|
||||
def _merge_openapi_tool_request_headers(
|
||||
static_headers: Dict[str, str]
|
||||
static_headers: Dict[str, str],
|
||||
) -> Dict[str, str]:
|
||||
"""Merge static closure headers with per-request ContextVar overrides.
|
||||
|
||||
|
|
|
|||
|
|
@ -1187,6 +1187,127 @@ async def get_end_user_object(
|
|||
return None
|
||||
|
||||
|
||||
_END_USER_VALIDATION_NEGATIVE_TTL = 60
|
||||
_END_USER_VALIDATION_POSITIVE_TTL = 300
|
||||
|
||||
|
||||
async def resolve_and_validate_end_user_id(
|
||||
raw_end_user_id: Optional[str],
|
||||
prisma_client: Optional[PrismaClient],
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
parent_otel_span: Optional[Span] = None,
|
||||
proxy_logging_obj: Optional[ProxyLogging] = None,
|
||||
route: str = "",
|
||||
) -> Optional[str]:
|
||||
"""Optionally drop end-user ids that don't resolve to a known DB row.
|
||||
|
||||
Default: pass-through. LiteLLM's documented pattern is that the `user`
|
||||
field is an arbitrary caller-supplied identifier, so validation is
|
||||
opt-in behind ``litellm.validate_end_user_id_in_db`` to preserve
|
||||
backwards compatibility.
|
||||
|
||||
When the flag is set: accept the id when it matches any of
|
||||
- LiteLLM_EndUserTable.user_id
|
||||
- LiteLLM_UserTable.user_id
|
||||
- LiteLLM_UserTable.user_email (case-insensitive)
|
||||
|
||||
If the id doesn't match but ``litellm.max_end_user_budget_id`` is set,
|
||||
we still preserve the id so the default end-user budget is applied
|
||||
downstream; otherwise we return None.
|
||||
|
||||
DB lookups reuse ``get_end_user_object`` / ``get_user_object`` so they
|
||||
share the same cache as the rest of the auth path instead of adding new
|
||||
raw Prisma queries.
|
||||
"""
|
||||
if raw_end_user_id is None:
|
||||
return None
|
||||
if not litellm.validate_end_user_id_in_db:
|
||||
return raw_end_user_id
|
||||
if prisma_client is None:
|
||||
return raw_end_user_id
|
||||
|
||||
cache_key = f"end_user_validation:{raw_end_user_id}"
|
||||
cached = await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
if cached == "valid":
|
||||
return raw_end_user_id
|
||||
if cached == "invalid":
|
||||
return raw_end_user_id if litellm.max_end_user_budget_id else None
|
||||
|
||||
is_valid = await _end_user_id_exists_in_db(
|
||||
end_user_id=raw_end_user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
route=route,
|
||||
)
|
||||
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=cache_key,
|
||||
value="valid" if is_valid else "invalid",
|
||||
ttl=(
|
||||
_END_USER_VALIDATION_POSITIVE_TTL
|
||||
if is_valid
|
||||
else _END_USER_VALIDATION_NEGATIVE_TTL
|
||||
),
|
||||
)
|
||||
|
||||
if is_valid:
|
||||
return raw_end_user_id
|
||||
# Preserve id so the caller can still apply litellm.max_end_user_budget_id.
|
||||
if litellm.max_end_user_budget_id:
|
||||
return raw_end_user_id
|
||||
return None
|
||||
|
||||
|
||||
async def _end_user_id_exists_in_db(
|
||||
end_user_id: str,
|
||||
prisma_client: PrismaClient,
|
||||
user_api_key_cache: UserApiKeyCache,
|
||||
parent_otel_span: Optional[Span] = None,
|
||||
proxy_logging_obj: Optional[ProxyLogging] = None,
|
||||
route: str = "",
|
||||
) -> bool:
|
||||
"""True when the id matches an EndUser, User, or user_email row."""
|
||||
try:
|
||||
end_user_obj = await get_end_user_object(
|
||||
end_user_id=end_user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
route=route,
|
||||
)
|
||||
if end_user_obj is not None:
|
||||
return True
|
||||
except litellm.BudgetExceededError:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
f"end_user validation: get_end_user_object lookup failed: {e}"
|
||||
)
|
||||
|
||||
try:
|
||||
user_obj = await get_user_object(
|
||||
user_id=end_user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_db_only=False,
|
||||
user_email=end_user_id if "@" in end_user_id else None,
|
||||
)
|
||||
if user_obj is not None:
|
||||
return True
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
f"end_user validation: get_user_object lookup failed: {e}"
|
||||
)
|
||||
|
||||
return False
|
||||
|
||||
|
||||
@log_db_metrics
|
||||
async def get_tag_objects_batch(
|
||||
tag_names: List[str],
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ import litellm
|
|||
from litellm import Router, provider_list
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import STANDARD_CUSTOMER_ID_HEADERS
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
from litellm.litellm_core_utils.url_utils import SSRFError, validate_url
|
||||
from litellm.proxy._types import *
|
||||
from litellm.types.router import CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS
|
||||
|
|
@ -1008,12 +1009,47 @@ def _get_customer_id_from_standard_headers(
|
|||
for standard_header in STANDARD_CUSTOMER_ID_HEADERS:
|
||||
for header_name, header_value in request_headers.items():
|
||||
if header_name.lower() == standard_header.lower():
|
||||
user_id_str = str(header_value) if header_value is not None else ""
|
||||
if user_id_str.strip():
|
||||
user_id_str = _coerce_user_id_to_str(header_value)
|
||||
if user_id_str:
|
||||
return user_id_str
|
||||
return None
|
||||
|
||||
|
||||
def _coerce_user_id_to_str(value: Any) -> Optional[str]:
|
||||
"""Return a usable end-user identifier string, or None if the value isn't one.
|
||||
|
||||
Always drops non-string structured values (dict/list/tuple/set) because
|
||||
stringifying them produces garbage spend-log rows like
|
||||
``"{'device_id': ...}"``. Strings that *decode* to a structured payload
|
||||
are only rejected when ``litellm.validate_end_user_id_in_db`` is enabled
|
||||
— operators who currently pass JSON-encoded identifiers keep their
|
||||
existing behavior until they opt in. See
|
||||
auth_utils.py:get_end_user_id_from_request_body for the extraction chain.
|
||||
"""
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, bool):
|
||||
# bool is an int subclass; handle explicitly to avoid "True"/"False".
|
||||
return None
|
||||
if isinstance(value, (int, float)):
|
||||
return str(value)
|
||||
if isinstance(value, str):
|
||||
stripped = value.strip()
|
||||
if not stripped:
|
||||
return None
|
||||
# Reject strings that decode to a structured payload (JSON object/array)
|
||||
# only when the operator has opted into end-user validation. Gating
|
||||
# behind the flag preserves backwards compatibility for deployments
|
||||
# that intentionally pass JSON-encoded user identifiers.
|
||||
if litellm.validate_end_user_id_in_db and stripped[:1] in ("{", "["):
|
||||
parsed = safe_json_loads(stripped)
|
||||
if isinstance(parsed, (dict, list)):
|
||||
return None
|
||||
return stripped
|
||||
# dict, list, tuple, set, arbitrary objects -> drop.
|
||||
return None
|
||||
|
||||
|
||||
def get_end_user_id_from_request_body(
|
||||
request_body: dict, request_headers: Optional[dict] = None
|
||||
) -> Optional[str]:
|
||||
|
|
@ -1052,23 +1088,22 @@ def get_end_user_id_from_request_body(
|
|||
if isinstance(custom_header_name_to_check, list):
|
||||
headers_lower = {k.lower(): v for k, v in request_headers.items()}
|
||||
for expected_header in custom_header_name_to_check:
|
||||
header_value = headers_lower.get(expected_header)
|
||||
if header_value is not None:
|
||||
user_id_str = str(header_value)
|
||||
if user_id_str.strip():
|
||||
return user_id_str
|
||||
user_id_str = _coerce_user_id_to_str(headers_lower.get(expected_header))
|
||||
if user_id_str:
|
||||
return user_id_str
|
||||
|
||||
elif isinstance(custom_header_name_to_check, str):
|
||||
for header_name, header_value in request_headers.items():
|
||||
if header_name.lower() == custom_header_name_to_check.lower():
|
||||
user_id_str = str(header_value) if header_value is not None else ""
|
||||
if user_id_str.strip():
|
||||
user_id_str = _coerce_user_id_to_str(header_value)
|
||||
if user_id_str:
|
||||
return user_id_str
|
||||
|
||||
# Check 3: 'user' field in request_body (commonly OpenAI)
|
||||
if "user" in request_body and request_body["user"] is not None:
|
||||
user_from_body_user_field = request_body["user"]
|
||||
return str(user_from_body_user_field)
|
||||
if "user" in request_body:
|
||||
user_id_str = _coerce_user_id_to_str(request_body["user"])
|
||||
if user_id_str:
|
||||
return user_id_str
|
||||
|
||||
def _as_dict(value: Any) -> dict:
|
||||
# metadata / litellm_metadata can arrive as JSON strings from
|
||||
|
|
@ -1077,32 +1112,30 @@ def get_end_user_id_from_request_body(
|
|||
if isinstance(value, dict):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
|
||||
parsed = safe_json_loads(value)
|
||||
return parsed if isinstance(parsed, dict) else {}
|
||||
return {}
|
||||
|
||||
# Check 4: 'litellm_metadata.user' in request_body (commonly Anthropic)
|
||||
litellm_metadata = _as_dict(request_body.get("litellm_metadata"))
|
||||
user_from_litellm_metadata = litellm_metadata.get("user")
|
||||
if user_from_litellm_metadata is not None:
|
||||
return str(user_from_litellm_metadata)
|
||||
user_id_str = _coerce_user_id_to_str(litellm_metadata.get("user"))
|
||||
if user_id_str:
|
||||
return user_id_str
|
||||
|
||||
# Check 5: 'metadata.user_id' in request_body (another common pattern)
|
||||
metadata_dict = _as_dict(request_body.get("metadata"))
|
||||
user_id_from_metadata_field = metadata_dict.get("user_id")
|
||||
if user_id_from_metadata_field is not None:
|
||||
return str(user_id_from_metadata_field)
|
||||
user_id_str = _coerce_user_id_to_str(metadata_dict.get("user_id"))
|
||||
if user_id_str:
|
||||
return user_id_str
|
||||
|
||||
# Check 6: 'safety_identifier' in request body (OpenAI Responses API parameter)
|
||||
# SECURITY NOTE: safety_identifier can be set by any caller in the request body.
|
||||
# Only use this for end-user identification in trusted environments where you control
|
||||
# the calling application. For untrusted callers, prefer using headers or server-side
|
||||
# middleware to set the end_user_id to prevent impersonation.
|
||||
if request_body.get("safety_identifier") is not None:
|
||||
user_from_body_user_field = request_body["safety_identifier"]
|
||||
return str(user_from_body_user_field)
|
||||
user_id_str = _coerce_user_id_to_str(request_body.get("safety_identifier"))
|
||||
if user_id_str:
|
||||
return user_id_str
|
||||
|
||||
return None
|
||||
|
||||
|
|
|
|||
|
|
@ -4,12 +4,15 @@ from typing import Dict, List, Optional, Set
|
|||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
|
||||
from litellm.proxy._types import SpecialModelNames, UserAPIKeyAuth
|
||||
from litellm.router import Router
|
||||
from litellm.router_utils.fallback_event_handlers import get_fallback_model_group
|
||||
from litellm.types.router import LiteLLM_Params
|
||||
from litellm.types.router import CredentialLiteLLMParams, LiteLLM_Params
|
||||
from litellm.utils import get_valid_models
|
||||
|
||||
_CREDENTIAL_LITELLM_PARAM_FIELDS = set(CredentialLiteLLMParams.model_fields)
|
||||
|
||||
|
||||
def _check_wildcard_routing(model: str) -> bool:
|
||||
"""
|
||||
|
|
@ -178,6 +181,7 @@ def get_complete_model_list(
|
|||
model_access_groups: Dict[str, List[str]] = {},
|
||||
include_model_access_groups: Optional[bool] = False,
|
||||
only_model_access_groups: Optional[bool] = False,
|
||||
team_id: Optional[str] = None,
|
||||
) -> List[str]:
|
||||
"""Logic for returning complete model list for a given key + team pair"""
|
||||
|
||||
|
|
@ -222,6 +226,7 @@ def get_complete_model_list(
|
|||
unique_models=unique_models,
|
||||
return_wildcard_routes=return_wildcard_routes,
|
||||
llm_router=llm_router,
|
||||
team_id=team_id,
|
||||
)
|
||||
|
||||
complete_model_list = unique_models + all_wildcard_models
|
||||
|
|
@ -229,6 +234,29 @@ def get_complete_model_list(
|
|||
return complete_model_list
|
||||
|
||||
|
||||
def _hydrate_litellm_credential_name(
|
||||
litellm_params: Optional[LiteLLM_Params],
|
||||
) -> Optional[LiteLLM_Params]:
|
||||
if litellm_params is None or litellm_params.litellm_credential_name is None:
|
||||
return litellm_params
|
||||
|
||||
credential_values = CredentialAccessor.get_credential_values(
|
||||
litellm_params.litellm_credential_name
|
||||
)
|
||||
if not credential_values:
|
||||
return litellm_params
|
||||
|
||||
litellm_params = litellm_params.model_copy()
|
||||
for key, value in credential_values.items():
|
||||
if (
|
||||
key in _CREDENTIAL_LITELLM_PARAM_FIELDS
|
||||
and getattr(litellm_params, key, None) is None
|
||||
):
|
||||
setattr(litellm_params, key, value)
|
||||
litellm_params.litellm_credential_name = None
|
||||
return litellm_params
|
||||
|
||||
|
||||
def get_known_models_from_wildcard(
|
||||
wildcard_model: str, litellm_params: Optional[LiteLLM_Params] = None
|
||||
) -> List[str]:
|
||||
|
|
@ -247,7 +275,7 @@ def get_known_models_from_wildcard(
|
|||
else:
|
||||
provider = wildcard_provider_prefix
|
||||
|
||||
# get all known provider models
|
||||
litellm_params = _hydrate_litellm_credential_name(litellm_params)
|
||||
|
||||
wildcard_models = get_provider_models(
|
||||
provider=provider, litellm_params=litellm_params
|
||||
|
|
@ -285,6 +313,7 @@ def _get_wildcard_models(
|
|||
unique_models: List[str],
|
||||
return_wildcard_routes: Optional[bool] = False,
|
||||
llm_router: Optional[Router] = None,
|
||||
team_id: Optional[str] = None,
|
||||
) -> List[str]:
|
||||
models_to_remove = set()
|
||||
all_wildcard_models = []
|
||||
|
|
@ -297,7 +326,9 @@ def _get_wildcard_models(
|
|||
|
||||
## get litellm params from model
|
||||
if llm_router is not None:
|
||||
model_list = llm_router.get_model_list(model_name=model)
|
||||
model_list = llm_router.get_model_list(
|
||||
model_name=model, team_id=team_id
|
||||
)
|
||||
if model_list:
|
||||
for router_model in model_list:
|
||||
wildcard_models = get_known_models_from_wildcard(
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ import fnmatch
|
|||
import re
|
||||
import secrets
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Iterator, List, NamedTuple, Optional, Tuple, Union, cast
|
||||
from typing import Any, Dict, Iterator, List, Optional, Tuple, Union, cast
|
||||
|
||||
import fastapi
|
||||
from fastapi import HTTPException, Request, WebSocket, status
|
||||
|
|
@ -44,6 +44,7 @@ from litellm.proxy.auth.auth_checks import (
|
|||
get_team_object,
|
||||
get_user_object,
|
||||
is_valid_fallback_model,
|
||||
resolve_and_validate_end_user_id,
|
||||
)
|
||||
from litellm.proxy.auth.auth_exception_handler import UserAPIKeyAuthExceptionHandler
|
||||
from litellm.proxy.auth.auth_utils import (
|
||||
|
|
@ -333,8 +334,22 @@ def _apply_budget_limits_to_end_user_params(
|
|||
async def user_api_key_auth_websocket(websocket: WebSocket):
|
||||
# Accept the WebSocket connection
|
||||
|
||||
scope_headers = list(websocket.scope.get("headers") or [])
|
||||
request = Request(scope={"type": "http", "headers": scope_headers})
|
||||
ws_scope = websocket.scope or {}
|
||||
scope_headers = list(ws_scope.get("headers") or [])
|
||||
# ``get_request_route`` falls back to ``request.url.path`` when
|
||||
# ``scope["path"]`` is absent. On WebSockets that fallback reads
|
||||
# ``websocket.url``, which Starlette reconstructs from the (poisonable)
|
||||
# Host header. Carry the ASGI scope's path / root_path so the lookup
|
||||
# never reaches the fallback.
|
||||
synthetic_scope: Dict[str, Any] = {
|
||||
"type": "http",
|
||||
"headers": scope_headers,
|
||||
"path": ws_scope.get("path", ""),
|
||||
}
|
||||
for key in ("root_path", "app_root_path"):
|
||||
if key in ws_scope:
|
||||
synthetic_scope[key] = ws_scope[key]
|
||||
request = Request(scope=synthetic_scope)
|
||||
|
||||
request._url = websocket.url
|
||||
|
||||
|
|
@ -1333,9 +1348,17 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
_end_user_object = None
|
||||
end_user_params = {}
|
||||
|
||||
end_user_id = get_end_user_id_from_request_body(
|
||||
raw_end_user_id = get_end_user_id_from_request_body(
|
||||
request_data, _safe_get_request_headers(request)
|
||||
)
|
||||
end_user_id = await resolve_and_validate_end_user_id(
|
||||
raw_end_user_id=raw_end_user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
route=route,
|
||||
)
|
||||
if end_user_id:
|
||||
try:
|
||||
end_user_params["end_user_id"] = end_user_id
|
||||
|
|
@ -2021,7 +2044,7 @@ def _team_obj_from_token(valid_token: UserAPIKeyAuth) -> LiteLLM_TeamTableCached
|
|||
|
||||
|
||||
@tracer.wrap()
|
||||
async def _run_centralized_common_checks(
|
||||
async def _run_centralized_common_checks( # noqa: PLR0915
|
||||
user_api_key_auth_obj: UserAPIKeyAuth,
|
||||
request: Request,
|
||||
request_data: dict,
|
||||
|
|
@ -2099,9 +2122,23 @@ async def _run_centralized_common_checks(
|
|||
return
|
||||
|
||||
parent_otel_span = user_api_key_auth_obj.parent_otel_span
|
||||
end_user_id = get_end_user_id_from_request_body(
|
||||
request_data, _safe_get_request_headers(request)
|
||||
)
|
||||
# In the integrated auth flow ``_user_api_key_auth_builder`` has already
|
||||
# resolved the end-user id and attached it here. Reuse that to avoid a
|
||||
# second extraction pass; fall back to extracting locally when the
|
||||
# function is invoked in isolation (e.g. in direct unit tests).
|
||||
end_user_id = user_api_key_auth_obj.end_user_id
|
||||
if end_user_id is None:
|
||||
raw_end_user_id = get_end_user_id_from_request_body(
|
||||
request_data, _safe_get_request_headers(request)
|
||||
)
|
||||
end_user_id = await resolve_and_validate_end_user_id(
|
||||
raw_end_user_id=raw_end_user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
route=route,
|
||||
)
|
||||
|
||||
fetch_coros = []
|
||||
if user_api_key_auth_obj.team_id is not None:
|
||||
|
|
@ -2432,11 +2469,33 @@ async def user_api_key_auth(
|
|||
api_key=api_key,
|
||||
)
|
||||
|
||||
end_user_id = get_end_user_id_from_request_body(
|
||||
request_data, _safe_get_request_headers(request)
|
||||
)
|
||||
if end_user_id is not None:
|
||||
user_api_key_auth_obj.end_user_id = end_user_id
|
||||
# Defense-in-depth: ``_user_api_key_auth_builder`` has multiple early-return
|
||||
# paths (no master key, /user/auth route, JWT short-circuits) that bypass
|
||||
# the end-user resolution block. If those paths produced an auth obj
|
||||
# without an ``end_user_id`` set, fall back to extracting from the request
|
||||
# body so spend logs are still attributed correctly. Validation honours
|
||||
# ``litellm.validate_end_user_id_in_db``.
|
||||
if user_api_key_auth_obj.end_user_id is None:
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
raw_end_user_id = get_end_user_id_from_request_body(
|
||||
request_data, _safe_get_request_headers(request)
|
||||
)
|
||||
if raw_end_user_id is not None:
|
||||
resolved_end_user_id = await resolve_and_validate_end_user_id(
|
||||
raw_end_user_id=raw_end_user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=user_api_key_auth_obj.parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
route=route,
|
||||
)
|
||||
if resolved_end_user_id is not None:
|
||||
user_api_key_auth_obj.end_user_id = resolved_end_user_id
|
||||
|
||||
user_api_key_auth_obj.request_route = normalize_request_route(route)
|
||||
return user_api_key_auth_obj
|
||||
|
|
|
|||
|
|
@ -324,7 +324,7 @@ class CustomOpenAPISpec:
|
|||
|
||||
@staticmethod
|
||||
def add_chat_completion_request_schema(
|
||||
openapi_schema: Dict[str, Any]
|
||||
openapi_schema: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Add ProxyChatCompletionRequest schema to chat completion endpoints for documentation.
|
||||
|
|
@ -380,7 +380,7 @@ class CustomOpenAPISpec:
|
|||
|
||||
@staticmethod
|
||||
def add_responses_api_request_schema(
|
||||
openapi_schema: Dict[str, Any]
|
||||
openapi_schema: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Add ResponsesAPIRequestParams schema to responses API endpoints for documentation.
|
||||
|
|
@ -410,7 +410,7 @@ class CustomOpenAPISpec:
|
|||
|
||||
@staticmethod
|
||||
def add_llm_api_request_schema_body(
|
||||
openapi_schema: Dict[str, Any]
|
||||
openapi_schema: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Add LLM API request schema bodies to OpenAPI specification for documentation.
|
||||
|
|
|
|||
|
|
@ -12,6 +12,33 @@ from litellm.proxy.common_utils.callback_utils import (
|
|||
)
|
||||
from litellm.types.router import Deployment
|
||||
|
||||
_FORM_CONTENT_TYPES: frozenset[str] = frozenset(
|
||||
{"application/x-www-form-urlencoded", "multipart/form-data"}
|
||||
)
|
||||
|
||||
|
||||
def _normalize_media_type(content_type: str) -> str:
|
||||
"""Return the bare media type per RFC 7231: strip params, trim, lowercase."""
|
||||
if not content_type:
|
||||
return ""
|
||||
return content_type.split(";", 1)[0].strip().lower()
|
||||
|
||||
|
||||
def _is_form_content_type(content_type: str) -> bool:
|
||||
"""
|
||||
True iff Starlette's ``request.form()`` will actually parse this body.
|
||||
|
||||
Substring matching ``"form"`` is unsafe: ``request.form()`` returns empty
|
||||
``FormData`` for non-canonical types without consuming the body, leaving
|
||||
the auth-time pre-read and the handler's read seeing different payloads.
|
||||
"""
|
||||
return _normalize_media_type(content_type) in _FORM_CONTENT_TYPES
|
||||
|
||||
|
||||
def _is_json_content_type(content_type: str) -> bool:
|
||||
"""True iff the body should be parsed as JSON."""
|
||||
return _normalize_media_type(content_type) == "application/json"
|
||||
|
||||
|
||||
async def _read_request_body(request: Optional[Request]) -> Dict:
|
||||
"""
|
||||
|
|
@ -37,8 +64,24 @@ async def _read_request_body(request: Optional[Request]) -> Dict:
|
|||
_request_headers: dict = _safe_get_request_headers(request=request)
|
||||
content_type = _request_headers.get("content-type", "")
|
||||
|
||||
if "form" in content_type:
|
||||
parsed_body = dict(await request.form())
|
||||
if _is_form_content_type(content_type):
|
||||
try:
|
||||
form_data = await request.form()
|
||||
except Exception as e:
|
||||
# ``request.form()`` raises on malformed multipart (missing
|
||||
# boundary, malformed chunk encoding, …). Surface as 400 so
|
||||
# the auth-time pre-read does not silently cache ``{}`` while
|
||||
# a later raw-body re-read sees the original payload —
|
||||
# banned-param checks must see the same body the handler
|
||||
# acts on.
|
||||
verbose_proxy_logger.error(f"Invalid form payload: {e}")
|
||||
raise ProxyException(
|
||||
message=f"Invalid form payload: {e}",
|
||||
type="invalid_request_error",
|
||||
param="request_body",
|
||||
code=status.HTTP_400_BAD_REQUEST,
|
||||
)
|
||||
parsed_body = dict(form_data)
|
||||
if "metadata" in parsed_body and isinstance(parsed_body["metadata"], str):
|
||||
parsed_body["metadata"] = json.loads(parsed_body["metadata"])
|
||||
else:
|
||||
|
|
@ -257,7 +300,7 @@ async def get_form_data(request: Request) -> Dict[str, Any]:
|
|||
|
||||
|
||||
async def convert_upload_files_to_file_data(
|
||||
form_data: Dict[str, Any]
|
||||
form_data: Dict[str, Any],
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Convert FastAPI UploadFile objects to file data tuples for litellm.
|
||||
|
|
@ -306,18 +349,13 @@ async def get_request_body(request: Request) -> Dict[str, Any]:
|
|||
Read the request body and parse it as JSON.
|
||||
"""
|
||||
if request.method == "POST":
|
||||
if request.headers.get("content-type", "") == "application/json":
|
||||
content_type = request.headers.get("content-type", "")
|
||||
if _is_json_content_type(content_type):
|
||||
return await _read_request_body(request)
|
||||
elif "multipart/form-data" in request.headers.get(
|
||||
"content-type", ""
|
||||
) or "application/x-www-form-urlencoded" in request.headers.get(
|
||||
"content-type", ""
|
||||
):
|
||||
elif _is_form_content_type(content_type):
|
||||
return await get_form_data(request)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported content type: {request.headers.get('content-type')}"
|
||||
)
|
||||
raise ValueError(f"Unsupported content type: {content_type}")
|
||||
return {}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Contains utils used by OpenAI compatible endpoints
|
||||
Contains utils used by OpenAI compatible endpoints
|
||||
"""
|
||||
|
||||
from typing import Optional, Set
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
What is this?
|
||||
What is this?
|
||||
|
||||
CRUD endpoints for managing pass-through endpoints
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -34,8 +34,7 @@ async def create_missing_views(db: _db): # noqa: PLR0915
|
|||
if not any(marker in error_msg for marker in _VIEW_NOT_FOUND_MARKERS):
|
||||
raise
|
||||
# If an error occurs, the view does not exist, so create it
|
||||
await db.execute_raw(
|
||||
"""
|
||||
await db.execute_raw("""
|
||||
CREATE VIEW "LiteLLM_VerificationTokenView" AS
|
||||
SELECT
|
||||
v.*,
|
||||
|
|
@ -47,8 +46,7 @@ async def create_missing_views(db: _db): # noqa: PLR0915
|
|||
FROM "LiteLLM_VerificationToken" v
|
||||
LEFT JOIN "LiteLLM_TeamTable" t ON v.team_id = t.team_id
|
||||
LEFT JOIN "LiteLLM_ProjectTable" p ON v.project_id = p.project_id;
|
||||
"""
|
||||
)
|
||||
""")
|
||||
|
||||
verbose_logger.debug("LiteLLM_VerificationTokenView Created!")
|
||||
|
||||
|
|
|
|||
|
|
@ -10,7 +10,6 @@ every text fragment.
|
|||
|
||||
from typing import Any, Callable, Dict, FrozenSet, Iterator, List
|
||||
|
||||
|
||||
# Call types whose body carries free-form chat / prompt text that
|
||||
# text-content guardrails (banned keywords, content moderation, secret
|
||||
# detection, …) should inspect. The proxy ingress passes ``route_type``
|
||||
|
|
|
|||
|
|
@ -4,7 +4,6 @@ from litellm.types.guardrails import SupportedGuardrailIntegrations
|
|||
|
||||
from .akto import AktoGuardrail
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.guardrails import Guardrail, LitellmParams
|
||||
|
||||
|
|
|
|||
|
|
@ -63,6 +63,7 @@ from litellm.types.utils import (
|
|||
CallTypesLiteral,
|
||||
Choices,
|
||||
GuardrailStatus,
|
||||
GuardrailTracingDetail,
|
||||
Message,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
|
|
@ -509,6 +510,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
# Add guardrail information to request trace
|
||||
#########################################################
|
||||
_json_response = httpx_response.json()
|
||||
tracing_detail = self._build_tracing_detail(_json_response)
|
||||
|
||||
# Raw Bedrock JSON is passed here; match/regex redaction runs once inside
|
||||
# CustomGuardrail.add_standard_logging_guardrail_information_to_request_data.
|
||||
self.add_standard_logging_guardrail_information_to_request_data(
|
||||
|
|
@ -522,6 +525,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
end_time=datetime.now().timestamp(),
|
||||
duration=(datetime.now() - start_time).total_seconds(),
|
||||
event_type=event_type,
|
||||
tracing_detail=tracing_detail or None,
|
||||
)
|
||||
#########################################################
|
||||
if httpx_response.status_code == 200:
|
||||
|
|
@ -640,6 +644,55 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
|
|||
return (status_code, err)
|
||||
return (status_code, message)
|
||||
|
||||
def _build_tracing_detail(
|
||||
self, response: BedrockGuardrailResponse
|
||||
) -> GuardrailTracingDetail:
|
||||
"""
|
||||
Build the tracing detail from the raw Bedrock response, before
|
||||
redaction, so downstream loggers (OTEL, Langfuse, ...) get the
|
||||
actual category names rather than the "[REDACTED]" sentinel that
|
||||
replaces customWords.match later. Bedrock's top-level ``action``
|
||||
field ("GUARDRAIL_INTERVENED" or "NONE") is also surfaced so the
|
||||
OTEL integration can expose it as a queryable span attribute
|
||||
without re-parsing the redacted guardrail_response blob.
|
||||
"""
|
||||
tracing_detail: GuardrailTracingDetail = {}
|
||||
violation_categories = self._extract_violation_category_names(response)
|
||||
if violation_categories:
|
||||
tracing_detail["violation_categories"] = violation_categories
|
||||
bedrock_action = response.get("action")
|
||||
if isinstance(bedrock_action, str):
|
||||
tracing_detail["guardrail_action"] = bedrock_action
|
||||
return tracing_detail
|
||||
|
||||
def _extract_violation_category_names(
|
||||
self, response: BedrockGuardrailResponse
|
||||
) -> List[str]:
|
||||
"""
|
||||
Flatten the BLOCKED assessments into a list of human-readable category
|
||||
names suitable for queryable OTEL / standard-logging attributes.
|
||||
|
||||
SECURITY: only emits the non-sensitive policy *label* (topic name,
|
||||
content-filter type, PII entity type, named-regex name). The raw
|
||||
``match`` field is intentionally NOT used — it carries the user's
|
||||
original input that triggered the rule (e.g. a credit-card number
|
||||
that hit a regex, or the literal custom word). Surfacing it to
|
||||
telemetry would re-introduce the sensitive content the guardrail
|
||||
was supposed to keep out. Entries that only have a ``match`` (bare
|
||||
customWords, unnamed regexes) are therefore skipped — operators
|
||||
can still see the count in ``_extract_blocked_assessments`` which
|
||||
feeds the HTTP error detail.
|
||||
"""
|
||||
names: List[str] = []
|
||||
for block in self._extract_blocked_assessments(response):
|
||||
for match in block.get("matches", []) or []:
|
||||
# Allow-list non-sensitive labels only. Never fall back to
|
||||
# `match.get("match")` — that's user-submitted content.
|
||||
label = match.get("name") or match.get("type")
|
||||
if isinstance(label, str) and label:
|
||||
names.append(label)
|
||||
return names
|
||||
|
||||
def _extract_blocked_assessments(
|
||||
self, response: BedrockGuardrailResponse
|
||||
) -> List[dict]:
|
||||
|
|
|
|||
35
litellm/proxy/guardrails/guardrail_hooks/rubrik/__init__.py
Normal file
35
litellm/proxy/guardrails/guardrail_hooks/rubrik/__init__.py
Normal file
|
|
@ -0,0 +1,35 @@
|
|||
"""Rubrik guardrail integration for LiteLLM."""
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from litellm.integrations.rubrik import RubrikLogger
|
||||
from litellm.types.guardrails import SupportedGuardrailIntegrations
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.guardrails import Guardrail, LitellmParams
|
||||
|
||||
|
||||
def initialize_guardrail(
|
||||
litellm_params: "LitellmParams", guardrail: "Guardrail"
|
||||
) -> RubrikLogger:
|
||||
import litellm
|
||||
|
||||
rubrik_callback = RubrikLogger(
|
||||
api_key=litellm_params.api_key,
|
||||
api_base=litellm_params.api_base,
|
||||
guardrail_name=guardrail.get("guardrail_name", ""),
|
||||
event_hook=litellm_params.mode,
|
||||
default_on=litellm_params.default_on,
|
||||
)
|
||||
|
||||
litellm.logging_callback_manager.add_litellm_callback(rubrik_callback)
|
||||
return rubrik_callback
|
||||
|
||||
|
||||
guardrail_initializer_registry = {
|
||||
SupportedGuardrailIntegrations.RUBRIK.value: initialize_guardrail,
|
||||
}
|
||||
|
||||
guardrail_class_registry = {
|
||||
SupportedGuardrailIntegrations.RUBRIK.value: RubrikLogger,
|
||||
}
|
||||
|
|
@ -6,7 +6,7 @@ The actual skill logic is in litellm/llms/litellm_proxy/skills/.
|
|||
|
||||
Usage:
|
||||
from litellm.proxy.hooks.litellm_skills import SkillsInjectionHook
|
||||
|
||||
|
||||
# Register hook in proxy
|
||||
litellm.callbacks.append(SkillsInjectionHook())
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
"""
|
||||
BUDGET MANAGEMENT
|
||||
|
||||
All /budget management endpoints
|
||||
All /budget management endpoints
|
||||
|
||||
/budget/new
|
||||
/budget/new
|
||||
/budget/info
|
||||
/budget/update
|
||||
/budget/delete
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
"""
|
||||
CUSTOMER MANAGEMENT
|
||||
|
||||
All /customer management endpoints
|
||||
All /customer management endpoints
|
||||
|
||||
/customer/new
|
||||
/customer/new
|
||||
/customer/info
|
||||
/customer/update
|
||||
/customer/delete
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue