mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-10 14:22:21 +00:00
Merge branch 'main' into fix/release-notes-v1-82-3-helicone-langfuse
This commit is contained in:
@@ -42,7 +42,7 @@ commands:
|
||||
"pydantic==2.11.0" "mcp==1.25.0" "requests-mock>=1.12.1" \
|
||||
"responses==0.25.7" "pytest-xdist==3.6.1" "pytest-timeout==2.2.0" \
|
||||
"pytest-cov==5.0.0" "semantic_router==0.1.10" "fastapi-offline==1.7.3" \
|
||||
"a2a"
|
||||
"a2a" "parameterized>=0.9.0"
|
||||
- setup_litellm_enterprise_pip
|
||||
- save_cache:
|
||||
paths:
|
||||
@@ -1115,7 +1115,7 @@ jobs:
|
||||
for dir in "${IGNORE_DIRS[@]}"; do
|
||||
IGNORE_ARGS="$IGNORE_ARGS --ignore=$dir"
|
||||
done
|
||||
python -m pytest -v tests/llm_translation $IGNORE_ARGS --junitxml=test-results/junit.xml --durations=20 -n 8 --timeout=120 --timeout_method=thread
|
||||
python -m pytest -v tests/llm_translation $IGNORE_ARGS --junitxml=test-results/junit.xml --durations=20 -n 8 --timeout=120 --timeout_method=thread --retries 2 --retry-delay 5
|
||||
no_output_timeout: 15m
|
||||
|
||||
# Store test results
|
||||
@@ -1331,7 +1331,7 @@ jobs:
|
||||
command: |
|
||||
pwd
|
||||
ls
|
||||
python -m pytest -vv tests/unified_google_tests --cov=litellm --cov-report=xml -x -s -v --junitxml=test-results/junit.xml --durations=5
|
||||
python -m pytest -vv tests/unified_google_tests --cov=litellm --cov-report=xml -x -s -v --junitxml=test-results/junit.xml --durations=5 --retries 3 --retry-delay 5
|
||||
no_output_timeout: 15m
|
||||
- run:
|
||||
name: Rename the coverage files
|
||||
|
||||
@@ -28,9 +28,12 @@ jobs:
|
||||
find . -type d -name "__pycache__" -exec rm -rf {} + || true
|
||||
find . -name "*.pyc" -delete || true
|
||||
|
||||
- name: Check poetry.lock is up to date
|
||||
run: |
|
||||
poetry check --lock || (echo "❌ poetry.lock is out of sync with pyproject.toml. Run 'poetry lock' locally and commit the result." && exit 1)
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
poetry lock
|
||||
poetry install --with dev
|
||||
|
||||
- name: Check Black formatting
|
||||
|
||||
@@ -163,6 +163,9 @@ run_grype_scans() {
|
||||
"CVE-2026-25639" # axios - full fix requires 1.x major version bump; pinned to >=0.30.2 to clear other axios CVEs, upgrade to 1.x in follow-up
|
||||
"CVE-2026-2297" # Python 3.13 SourcelessFileLoader audit hook bypass - no fix available in base image
|
||||
"GHSA-qffp-2rhf-9h96" # tar hardlink path traversal - from nodejs_wheel bundled npm, not used in application runtime code
|
||||
"CVE-2026-2673" # OpenSSL 3.6.1 TLS 1.3 key exchange group negotiation issue - no fix available yet
|
||||
"CVE-2026-3644" # Python 3.13 vulnerability - no fix available in base image
|
||||
"CVE-2026-4224" # Python 3.13 Expat parser stack overflow in ElementDeclHandler - no fix available in base image
|
||||
)
|
||||
|
||||
# Build JSON array of allowlisted CVE IDs for jq
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
---
|
||||
slug: guardrail-logging-secret-exposure-incident
|
||||
title: "Incident Report: Guardrail logging exposed secret headers in spend logs and traces"
|
||||
date: 2026-03-18T10:00:00
|
||||
authors:
|
||||
- litellm
|
||||
tags: [incident-report, security, guardrails]
|
||||
hide_table_of_contents: false
|
||||
---
|
||||
|
||||
**Date:** March 18, 2026
|
||||
**Duration:** Unknown
|
||||
**Severity:** High
|
||||
**Status:** Resolved
|
||||
|
||||
## Summary
|
||||
|
||||
When a custom guardrail returned the full LiteLLM request/data dictionary, the guardrail response logged by LiteLLM could include `secret_fields.raw_headers`, including plaintext `Authorization` headers containing API keys or other credentials.
|
||||
|
||||
This information could then propagate to logging and observability surfaces that consume guardrail metadata, including:
|
||||
|
||||
- **Spend logs in the LiteLLM UI:** visible to admins with access to spend-log data
|
||||
- **OpenTelemetry traces:** visible to anyone with access to the relevant telemetry backend
|
||||
|
||||
LLM calls, proxy routing, and provider execution were not blocked by this bug. The impact was exposure of sensitive request headers in observability and logging paths.
|
||||
|
||||
{/* truncate */}
|
||||
|
||||
---
|
||||
|
||||
## Background
|
||||
|
||||
LiteLLM keeps internal request data (including request headers) for use during the call. That data is not meant to be written to logs or telemetry.
|
||||
|
||||
When custom guardrails run, their outcomes are logged so they can appear in spend logs, OpenTelemetry traces, and other observability backends. If a guardrail returned the full request payload instead of a minimal result, that internal request data could be included in what was logged. Before the fix, the guardrail logging path did not strip that data before sending it to those systems.
|
||||
|
||||
```mermaid
|
||||
flowchart TD
|
||||
inboundRequest["1. Incoming proxy request"] --> storeSecrets["2. Store internal request data"]
|
||||
storeSecrets --> guardrailRuns["3. Custom guardrail runs"]
|
||||
guardrailRuns --> fullDataReturn["4. Guardrail returns full request payload"]
|
||||
fullDataReturn --> loggingBuild["5. Build guardrail log payload"]
|
||||
loggingBuild --> spendLogs["6a. Persist to spend logs / UI"]
|
||||
loggingBuild --> otelTraces["6b. Attach to OTEL guardrail spans"]
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Root Cause
|
||||
|
||||
The root cause was incomplete sanitization in the guardrail logging path. When building the payload that gets sent to spend logs and traces, LiteLLM prepared guardrail responses for logging but did not strip internal request data (such as headers) from them. If a guardrail returned a response that included that data, it was passed through to the logging and observability systems unchanged.
|
||||
|
||||
---
|
||||
|
||||
## Impact
|
||||
|
||||
This issue required all of the following:
|
||||
|
||||
1. A custom guardrail returned the full LiteLLM request/data dictionary, or another response object containing `secret_fields`.
|
||||
2. LiteLLM logged that guardrail response through the standard guardrail logging path.
|
||||
3. An operator, admin, or telemetry consumer had access to the resulting logs or traces.
|
||||
|
||||
When those conditions were met, sensitive values could become visible through:
|
||||
|
||||
- **Spend logs / UI responses:** guardrail metadata could be included in spend-log payloads rendered in the admin UI.
|
||||
- **OpenTelemetry traces:** `guardrail_response` could be written as a span attribute on guardrail spans.
|
||||
- **Other downstream observability backends:** any integration consuming the same guardrail metadata could receive the leaked values.
|
||||
|
||||
This was a logging and telemetry exposure bug. It did not let callers bypass auth, access other tenants directly, or change model behavior, but it could expose plaintext credentials to people with access to those observability systems.
|
||||
|
||||
---
|
||||
|
||||
## Guidance For Users
|
||||
|
||||
- Upgrade to LiteLLM 1.82.3+.
|
||||
- If you operated custom guardrails that return the full request/data dict, review whether spend logs or telemetry traces were retained during the affected period.
|
||||
- Rotate any credentials that may have appeared in `Authorization` or other forwarded request headers in those systems.
|
||||
- Apply least-privilege access controls to spend-log views and telemetry backends that may contain request-derived metadata.
|
||||
@@ -902,6 +902,7 @@ router_settings:
|
||||
| OTEL_SERVICE_NAME | Service name identifier for OpenTelemetry
|
||||
| OTEL_TRACER_NAME | Tracer name for OpenTelemetry tracing
|
||||
| OTEL_LOGS_EXPORTER | Exporter type for OpenTelemetry logs (e.g., console)
|
||||
| OTEL_IGNORE_CONTEXT_PROPAGATION | When true, ignore parent span context propagation in OpenTelemetry callbacks
|
||||
| PAGERDUTY_API_KEY | API key for PagerDuty Alerting
|
||||
| PANW_PRISMA_AIRS_API_KEY | API key for PANW Prisma AIRS service
|
||||
| PANW_PRISMA_AIRS_API_BASE | Base URL for PANW Prisma AIRS service
|
||||
|
||||
@@ -602,6 +602,22 @@ Since you shouldn't use 12.5, round down to **10** to leave a safety buffer. Thi
|
||||
- Total maximum connections: 8 workers × 10 connections = 80 connections
|
||||
- This stays safely under your database's 100 connection limit
|
||||
|
||||
## LiteLLM License Key (Enterprise)
|
||||
|
||||
To enable [LiteLLM Enterprise features](https://docs.litellm.ai/docs/proxy/enterprise), set your license key as an environment variable:
|
||||
|
||||
```bash
|
||||
export LITELLM_LICENSE="eyJ..."
|
||||
```
|
||||
|
||||
The license key is a JWT token provided when you purchase a LiteLLM Enterprise license. Once set, LiteLLM will automatically detect and activate enterprise features.
|
||||
|
||||
You can also add it to your `.env` file:
|
||||
|
||||
```env
|
||||
LITELLM_LICENSE="eyJ..."
|
||||
```
|
||||
|
||||
## Extras
|
||||
|
||||
|
||||
|
||||
@@ -48,6 +48,20 @@ const sidebars = {
|
||||
slug: "/guardrail_providers"
|
||||
},
|
||||
items: [
|
||||
{
|
||||
type: "category",
|
||||
label: "Contributing to Guardrails",
|
||||
items: [
|
||||
"adding_provider/generic_guardrail_api",
|
||||
"adding_provider/simple_guardrail_tutorial",
|
||||
"adding_provider/adding_guardrail_support",
|
||||
]
|
||||
},
|
||||
{
|
||||
type: "doc",
|
||||
id: "proxy/guardrails/team_based_guardrails",
|
||||
label: "Team Bring-Your-Own Guardrails",
|
||||
},
|
||||
...[
|
||||
"proxy/guardrails/qualifire",
|
||||
"proxy/guardrails/aim_security",
|
||||
|
||||
@@ -757,7 +757,7 @@ def _map_traffic_type_to_service_tier(traffic_type: Optional[str]) -> Optional[s
|
||||
"""
|
||||
if traffic_type is None:
|
||||
return None
|
||||
service_tier = _GEMINI_TRAFFIC_TYPE_TO_SERVICE_TIER.get(traffic_type.upper())
|
||||
service_tier = _GEMINI_TRAFFIC_TYPE_TO_SERVICE_TIER.get(str(traffic_type).upper())
|
||||
return service_tier
|
||||
|
||||
|
||||
|
||||
@@ -291,7 +291,7 @@ class DataDogLogger(
|
||||
|
||||
dd_payload = DatadogPayload(
|
||||
ddsource=get_datadog_source(),
|
||||
ddtags=get_datadog_tags(),
|
||||
ddtags=",".join(get_datadog_tags()),
|
||||
hostname=get_datadog_hostname(),
|
||||
message=safe_dumps(message_payload),
|
||||
service=get_datadog_service(),
|
||||
@@ -442,7 +442,9 @@ class DataDogLogger(
|
||||
verbose_logger.debug("Datadog: Logger - Logging payload = %s", json_payload)
|
||||
dd_payload = DatadogPayload(
|
||||
ddsource=get_datadog_source(),
|
||||
ddtags=get_datadog_tags(standard_logging_object=standard_logging_object),
|
||||
ddtags=",".join(
|
||||
get_datadog_tags(standard_logging_object=standard_logging_object)
|
||||
),
|
||||
hostname=get_datadog_hostname(),
|
||||
message=json_payload,
|
||||
service=get_datadog_service(),
|
||||
@@ -545,7 +547,7 @@ class DataDogLogger(
|
||||
_dd_message_str = safe_dumps(_payload_dict)
|
||||
_dd_payload = DatadogPayload(
|
||||
ddsource=get_datadog_source(),
|
||||
ddtags=get_datadog_tags(),
|
||||
ddtags=",".join(get_datadog_tags()),
|
||||
hostname=get_datadog_hostname(),
|
||||
message=_dd_message_str,
|
||||
service=get_datadog_service(),
|
||||
@@ -587,7 +589,7 @@ class DataDogLogger(
|
||||
_dd_message_str = safe_dumps(_payload_dict)
|
||||
_dd_payload = DatadogPayload(
|
||||
ddsource=get_datadog_source(),
|
||||
ddtags=get_datadog_tags(),
|
||||
ddtags=",".join(get_datadog_tags()),
|
||||
hostname=get_datadog_hostname(),
|
||||
message=_dd_message_str,
|
||||
service=get_datadog_service(),
|
||||
@@ -678,7 +680,7 @@ class DataDogLogger(
|
||||
|
||||
dd_payload = DatadogPayload(
|
||||
ddsource=get_datadog_source(),
|
||||
ddtags=get_datadog_tags(),
|
||||
ddtags=",".join(get_datadog_tags()),
|
||||
hostname=get_datadog_hostname(),
|
||||
message=json_payload,
|
||||
service=get_datadog_service(),
|
||||
|
||||
@@ -38,8 +38,13 @@ def get_datadog_pod_name() -> str:
|
||||
|
||||
def get_datadog_tags(
|
||||
standard_logging_object: Optional[StandardLoggingPayload] = None,
|
||||
) -> str:
|
||||
"""Build Datadog tags string used by multiple integrations."""
|
||||
) -> List[str]:
|
||||
"""Build Datadog tags as a list of individual tag strings.
|
||||
|
||||
Returns a list of "key:value" strings suitable for Datadog LLM Observability
|
||||
(which expects tags as an array). For Datadog Logs API (ddtags), join with
|
||||
comma: ",".join(get_datadog_tags(...)).
|
||||
"""
|
||||
|
||||
base_tags = {
|
||||
"env": get_datadog_env(),
|
||||
@@ -66,4 +71,4 @@ def get_datadog_tags(
|
||||
if team_tag:
|
||||
tags.append(f"team:{team_tag}")
|
||||
|
||||
return ",".join(tags)
|
||||
return tags
|
||||
|
||||
@@ -203,7 +203,7 @@ class DataDogLLMObsLogger(CustomBatchLogger):
|
||||
type="span",
|
||||
attributes=DDSpanAttributes(
|
||||
ml_app=get_datadog_service(),
|
||||
tags=[get_datadog_tags()],
|
||||
tags=get_datadog_tags(),
|
||||
spans=self.log_queue,
|
||||
),
|
||||
),
|
||||
@@ -315,7 +315,7 @@ class DataDogLLMObsLogger(CustomBatchLogger):
|
||||
duration=int((end_time - start_time).total_seconds() * 1e9),
|
||||
metrics=metrics,
|
||||
status="error" if error_info else "ok",
|
||||
tags=[get_datadog_tags(standard_logging_object=standard_logging_payload)],
|
||||
tags=get_datadog_tags(standard_logging_object=standard_logging_payload),
|
||||
)
|
||||
|
||||
apm_trace_id = self._get_apm_trace_id()
|
||||
|
||||
@@ -5,7 +5,6 @@ import os
|
||||
import random
|
||||
import traceback
|
||||
import types
|
||||
from litellm._uuid import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
@@ -14,10 +13,11 @@ from pydantic import BaseModel # type: ignore
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.integrations.custom_batch_logger import CustomBatchLogger
|
||||
from litellm.integrations.langsmith_mock_client import (
|
||||
should_use_langsmith_mock,
|
||||
create_mock_langsmith_client,
|
||||
should_use_langsmith_mock,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
@@ -110,6 +110,60 @@ class LangsmithLogger(CustomBatchLogger):
|
||||
LANGSMITH_TENANT_ID=_credentials_tenant_id,
|
||||
)
|
||||
|
||||
def _extract_metadata_fields(
|
||||
self, metadata: dict, credentials: LangsmithCredentialsObject
|
||||
):
|
||||
return {
|
||||
"project_name": metadata.get(
|
||||
"project_name", credentials["LANGSMITH_PROJECT"]
|
||||
),
|
||||
"run_name": metadata.get("run_name", self.langsmith_default_run_name),
|
||||
"run_id": metadata.get("id", metadata.get("run_id", None)),
|
||||
"parent_run_id": metadata.get("parent_run_id", None),
|
||||
"trace_id": metadata.get("trace_id", None),
|
||||
"session_id": metadata.get("session_id", None),
|
||||
"dotted_order": metadata.get("dotted_order", None),
|
||||
}
|
||||
|
||||
def _build_extra_metadata(self, metadata: Dict):
|
||||
extra_metadata = dict(metadata)
|
||||
requester_metadata = extra_metadata.get("requester_metadata")
|
||||
if requester_metadata and isinstance(requester_metadata, dict):
|
||||
for key in ("session_id", "thread_id", "conversation_id"):
|
||||
if key in requester_metadata and key not in extra_metadata:
|
||||
extra_metadata[key] = requester_metadata[key]
|
||||
return extra_metadata
|
||||
|
||||
def _build_outputs_with_usage(
|
||||
self, payload: StandardLoggingPayload
|
||||
) -> Dict[str, Any]:
|
||||
response = payload["response"]
|
||||
outputs: Dict[str, Any]
|
||||
if isinstance(response, dict):
|
||||
outputs = {**response}
|
||||
else:
|
||||
outputs = {"output": response}
|
||||
outputs["usage_metadata"] = {
|
||||
"input_tokens": payload.get("prompt_tokens", 0),
|
||||
"output_tokens": payload.get("completion_tokens", 0),
|
||||
"total_tokens": payload.get("total_tokens", 0),
|
||||
"total_cost": payload.get("response_cost", 0),
|
||||
}
|
||||
return outputs
|
||||
|
||||
def _ensure_required_ids(self, data: dict, run_id: Optional[str]):
|
||||
if "id" not in data or data["id"] is None:
|
||||
run_id = str(uuid.uuid4())
|
||||
data["id"] = run_id
|
||||
|
||||
if "trace_id" not in data or data["trace_id"] is None:
|
||||
if run_id is not None and isinstance(run_id, str):
|
||||
data["trace_id"] = run_id
|
||||
|
||||
if "dotted_order" not in data or data["dotted_order"] is None:
|
||||
if run_id is not None and isinstance(run_id, str):
|
||||
data["dotted_order"] = self.make_dot_order(run_id=run_id)
|
||||
|
||||
def _prepare_log_data(
|
||||
self,
|
||||
kwargs,
|
||||
@@ -121,44 +175,28 @@ class LangsmithLogger(CustomBatchLogger):
|
||||
try:
|
||||
_litellm_params = kwargs.get("litellm_params", {}) or {}
|
||||
metadata = _litellm_params.get("metadata", {}) or {}
|
||||
project_name = metadata.get(
|
||||
"project_name", credentials["LANGSMITH_PROJECT"]
|
||||
)
|
||||
run_name = metadata.get("run_name", self.langsmith_default_run_name)
|
||||
run_id = metadata.get("id", metadata.get("run_id", None))
|
||||
parent_run_id = metadata.get("parent_run_id", None)
|
||||
trace_id = metadata.get("trace_id", None)
|
||||
session_id = metadata.get("session_id", None)
|
||||
dotted_order = metadata.get("dotted_order", None)
|
||||
|
||||
fields = self._extract_metadata_fields(metadata, credentials)
|
||||
verbose_logger.debug(
|
||||
f"Langsmith Logging - project_name: {project_name}, run_name {run_name}"
|
||||
f"Langsmith Logging - project_name: {fields['project_name']}, run_name {fields['run_name']}"
|
||||
)
|
||||
|
||||
# Ensure everything in the payload is converted to str
|
||||
payload: Optional[StandardLoggingPayload] = kwargs.get(
|
||||
"standard_logging_object", None
|
||||
)
|
||||
|
||||
if payload is None:
|
||||
raise Exception("Error logging request payload. Payload=none.")
|
||||
|
||||
metadata = payload[
|
||||
"metadata"
|
||||
] # ensure logged metadata is json serializable
|
||||
|
||||
extra_metadata = dict(metadata)
|
||||
requester_metadata = extra_metadata.get("requester_metadata")
|
||||
if requester_metadata and isinstance(requester_metadata, dict):
|
||||
for key in ("session_id", "thread_id", "conversation_id"):
|
||||
if key in requester_metadata and key not in extra_metadata:
|
||||
extra_metadata[key] = requester_metadata[key]
|
||||
metadata = payload["metadata"]
|
||||
extra_metadata = self._build_extra_metadata(dict(metadata))
|
||||
outputs = self._build_outputs_with_usage(payload)
|
||||
|
||||
data = {
|
||||
"name": run_name,
|
||||
"run_type": "llm", # this should always be llm, since litellm always logs llm calls. Langsmith allow us to log "chain"
|
||||
"name": fields["run_name"],
|
||||
"run_type": "llm",
|
||||
"inputs": payload,
|
||||
"outputs": payload["response"],
|
||||
"session_name": project_name,
|
||||
"outputs": outputs,
|
||||
"session_name": fields["project_name"],
|
||||
"start_time": payload["startTime"],
|
||||
"end_time": payload["endTime"],
|
||||
"tags": payload["request_tags"],
|
||||
@@ -168,46 +206,19 @@ class LangsmithLogger(CustomBatchLogger):
|
||||
if payload["error_str"] is not None and payload["status"] == "failure":
|
||||
data["error"] = payload["error_str"]
|
||||
|
||||
if run_id:
|
||||
data["id"] = run_id
|
||||
|
||||
if parent_run_id:
|
||||
data["parent_run_id"] = parent_run_id
|
||||
|
||||
if trace_id:
|
||||
data["trace_id"] = trace_id
|
||||
|
||||
if session_id:
|
||||
data["session_id"] = session_id
|
||||
|
||||
if dotted_order:
|
||||
data["dotted_order"] = dotted_order
|
||||
|
||||
run_id: Optional[str] = data.get("id") # type: ignore
|
||||
if "id" not in data or data["id"] is None:
|
||||
"""
|
||||
for /batch langsmith requires id, trace_id and dotted_order passed as params
|
||||
"""
|
||||
run_id = str(uuid.uuid4())
|
||||
|
||||
data["id"] = run_id
|
||||
|
||||
if (
|
||||
"trace_id" not in data
|
||||
or data["trace_id"] is None
|
||||
and (run_id is not None and isinstance(run_id, str))
|
||||
for key in (
|
||||
"id",
|
||||
"parent_run_id",
|
||||
"trace_id",
|
||||
"session_id",
|
||||
"dotted_order",
|
||||
):
|
||||
data["trace_id"] = run_id
|
||||
|
||||
if (
|
||||
"dotted_order" not in data
|
||||
or data["dotted_order"] is None
|
||||
and (run_id is not None and isinstance(run_id, str))
|
||||
):
|
||||
data["dotted_order"] = self.make_dot_order(run_id=run_id) # type: ignore
|
||||
field_key = "run_id" if key == "id" else key
|
||||
if fields[field_key]:
|
||||
data[key] = fields[field_key]
|
||||
|
||||
self._ensure_required_ids(data, fields["run_id"])
|
||||
verbose_logger.debug("Langsmith Logging data on langsmith: %s", data)
|
||||
|
||||
return data
|
||||
except Exception:
|
||||
raise
|
||||
|
||||
@@ -84,6 +84,8 @@ from litellm.types.llms.openai import (
|
||||
OpenAIModerationResponse,
|
||||
ResponseAPIUsage,
|
||||
ResponseCompletedEvent,
|
||||
ResponseFailedEvent,
|
||||
ResponseIncompleteEvent,
|
||||
ResponsesAPIResponse,
|
||||
)
|
||||
from litellm.types.mcp import MCPPostCallResponseObject
|
||||
@@ -516,6 +518,23 @@ class Logging(LiteLLMLoggingBaseClass):
|
||||
),
|
||||
)
|
||||
|
||||
def get_router_model_id(self) -> Optional[str]:
|
||||
"""Extract the router deployment model_id from litellm_params.
|
||||
|
||||
Checks both litellm_metadata and metadata for model_info.id.
|
||||
Used by cost calculators to look up custom pricing registered
|
||||
under the deployment's model_info.id in litellm.model_cost.
|
||||
"""
|
||||
if not hasattr(self, "litellm_params"):
|
||||
return None
|
||||
for key in ("litellm_metadata", "metadata"):
|
||||
meta = self.litellm_params.get(key, {}) or {}
|
||||
info = meta.get("model_info", {}) or {}
|
||||
model_id = info.get("id")
|
||||
if model_id is not None:
|
||||
return model_id
|
||||
return None
|
||||
|
||||
def update_environment_variables(
|
||||
self,
|
||||
litellm_params: Dict,
|
||||
@@ -1458,16 +1477,8 @@ class Logging(LiteLLMLoggingBaseClass):
|
||||
# Fallback: extract router_model_id from litellm_params when not available
|
||||
# from the result object. ResponsesAPIResponse objects (used by /v1/responses
|
||||
# streaming) don't carry _hidden_params["model_id"] like ModelResponse does.
|
||||
if router_model_id is None and hasattr(self, "litellm_params"):
|
||||
for metadata_key in ("litellm_metadata", "metadata"):
|
||||
_metadata: dict = (
|
||||
self.litellm_params.get(metadata_key, {}) or {}
|
||||
)
|
||||
_model_info: dict = _metadata.get("model_info", {}) or {}
|
||||
_model_id = _model_info.get("id")
|
||||
if _model_id is not None:
|
||||
router_model_id = _model_id
|
||||
break
|
||||
if router_model_id is None:
|
||||
router_model_id = self.get_router_model_id()
|
||||
|
||||
## RESPONSE COST ##
|
||||
custom_pricing = use_custom_pricing_for_model(
|
||||
@@ -2972,8 +2983,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
||||
if (
|
||||
isinstance(callback, CustomLogger)
|
||||
and is_sync_request
|
||||
and self.call_type
|
||||
!= CallTypes.pass_through.value
|
||||
and self.call_type != CallTypes.pass_through.value
|
||||
): # custom logger class
|
||||
callback.log_failure_event(
|
||||
start_time=start_time,
|
||||
@@ -3321,7 +3331,10 @@ class Logging(LiteLLMLoggingBaseClass):
|
||||
return result
|
||||
elif isinstance(result, TextCompletionResponse):
|
||||
return result
|
||||
elif isinstance(result, ResponseCompletedEvent):
|
||||
elif isinstance(
|
||||
result,
|
||||
(ResponseCompletedEvent, ResponseIncompleteEvent, ResponseFailedEvent),
|
||||
):
|
||||
## return unified Usage object
|
||||
if isinstance(result.response.usage, ResponseAPIUsage):
|
||||
transformed_usage = (
|
||||
@@ -3342,7 +3355,6 @@ class Logging(LiteLLMLoggingBaseClass):
|
||||
return result.response
|
||||
else:
|
||||
return None
|
||||
return None
|
||||
|
||||
def _handle_anthropic_messages_response_logging(self, result: Any) -> ModelResponse:
|
||||
"""
|
||||
|
||||
@@ -160,6 +160,7 @@ class CustomStreamWrapper:
|
||||
self.chunks: List = (
|
||||
[]
|
||||
) # keep track of the returned chunks - used for calculating the input/output tokens for stream options
|
||||
self._repeated_messages_count = 1
|
||||
self.is_function_call = self.check_is_function_call(logging_obj=logging_obj)
|
||||
self.created: Optional[int] = None
|
||||
self._last_returned_hidden_params: Optional[dict] = None
|
||||
@@ -241,7 +242,7 @@ class CustomStreamWrapper:
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
def safety_checker(self) -> None:
|
||||
def raise_on_model_repetition(self) -> None:
|
||||
"""
|
||||
Fixes - https://github.com/BerriAI/litellm/issues/5158
|
||||
|
||||
@@ -249,28 +250,35 @@ class CustomStreamWrapper:
|
||||
|
||||
Raises - InternalServerError, if LLM enters infinite loop while streaming
|
||||
"""
|
||||
if len(self.chunks) >= litellm.REPEATED_STREAMING_CHUNK_LIMIT:
|
||||
# Get the last n chunks
|
||||
last_chunks = self.chunks[-litellm.REPEATED_STREAMING_CHUNK_LIMIT :]
|
||||
if len(self.chunks) < 2:
|
||||
return
|
||||
|
||||
# Extract the relevant content from the chunks
|
||||
last_contents = [chunk.choices[0].delta.content for chunk in last_chunks]
|
||||
last_content = self.chunks[-1].choices[0].delta.content
|
||||
|
||||
# Check if all extracted contents are identical
|
||||
if all(content == last_contents[0] for content in last_contents):
|
||||
if (
|
||||
last_contents[0] is not None
|
||||
and isinstance(last_contents[0], str)
|
||||
and len(last_contents[0]) > 2
|
||||
): # ignore empty content - https://github.com/BerriAI/litellm/issues/5158#issuecomment-2287156946
|
||||
# All last n chunks are identical
|
||||
raise litellm.InternalServerError(
|
||||
message="The model is repeating the same chunk = {}.".format(
|
||||
last_contents[0]
|
||||
),
|
||||
model="",
|
||||
llm_provider="",
|
||||
)
|
||||
if (
|
||||
last_content is None
|
||||
or not isinstance(last_content, str)
|
||||
or len(last_content) <= 2
|
||||
): # ignore empty content - https://github.com/BerriAI/litellm/issues/5158#issuecomment-2287156946
|
||||
self._repeated_messages_count = 1
|
||||
return
|
||||
|
||||
second_to_last_content = self.chunks[-2].choices[0].delta.content
|
||||
|
||||
if last_content == second_to_last_content:
|
||||
self._repeated_messages_count += 1
|
||||
else:
|
||||
self._repeated_messages_count = 1
|
||||
|
||||
if self._repeated_messages_count >= litellm.REPEATED_STREAMING_CHUNK_LIMIT:
|
||||
# All last n chunks are identical
|
||||
raise litellm.InternalServerError(
|
||||
message="The model is repeating the same chunk = {}.".format(
|
||||
last_content
|
||||
),
|
||||
model="",
|
||||
llm_provider="",
|
||||
)
|
||||
|
||||
def check_special_tokens(self, chunk: str, finish_reason: Optional[str]):
|
||||
"""
|
||||
@@ -924,7 +932,7 @@ class CustomStreamWrapper:
|
||||
if (
|
||||
is_chunk_non_empty
|
||||
): # cannot set content of an OpenAI Object to be an empty string
|
||||
self.safety_checker()
|
||||
self.raise_on_model_repetition()
|
||||
hold, model_response_str = self.check_special_tokens(
|
||||
chunk=completion_obj["content"],
|
||||
finish_reason=model_response.choices[0].finish_reason,
|
||||
@@ -1893,15 +1901,19 @@ class CustomStreamWrapper:
|
||||
"usage",
|
||||
getattr(complete_streaming_response, "usage"),
|
||||
)
|
||||
try:
|
||||
_cache_copy = complete_streaming_response.model_copy(deep=True)
|
||||
_log_copy = complete_streaming_response.model_copy(deep=True)
|
||||
except RuntimeError:
|
||||
_cache_copy = complete_streaming_response.model_copy()
|
||||
_log_copy = complete_streaming_response.model_copy()
|
||||
self.cache_streaming_response(
|
||||
processed_chunk=complete_streaming_response.model_copy(
|
||||
deep=True
|
||||
),
|
||||
processed_chunk=_cache_copy,
|
||||
cache_hit=cache_hit,
|
||||
)
|
||||
executor.submit(
|
||||
self.logging_obj.success_handler,
|
||||
complete_streaming_response.model_copy(deep=True),
|
||||
_log_copy,
|
||||
None,
|
||||
None,
|
||||
cache_hit,
|
||||
@@ -2113,11 +2125,13 @@ class CustomStreamWrapper:
|
||||
"usage",
|
||||
getattr(complete_streaming_response, "usage"),
|
||||
)
|
||||
try:
|
||||
_copy = complete_streaming_response.model_copy(deep=True)
|
||||
except RuntimeError:
|
||||
_copy = complete_streaming_response.model_copy()
|
||||
asyncio.create_task(
|
||||
self.async_cache_streaming_response(
|
||||
processed_chunk=complete_streaming_response.model_copy(
|
||||
deep=True
|
||||
),
|
||||
processed_chunk=_copy,
|
||||
cache_hit=cache_hit,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -48,6 +48,10 @@ from litellm.types.llms.openai import (
|
||||
ChatCompletionToolCallChunk,
|
||||
ChatCompletionToolCallFunctionChunk,
|
||||
)
|
||||
from litellm.types.responses.main import (
|
||||
OutputCodeInterpreterCall,
|
||||
build_code_interpreter_log_outputs,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
Delta,
|
||||
GenericStreamingChunk,
|
||||
@@ -538,6 +542,12 @@ class ModelResponseIterator:
|
||||
# Accumulate compaction blocks for multi-turn reconstruction
|
||||
self.compaction_blocks: List[Dict[str, Any]] = []
|
||||
|
||||
# Track server tool use inputs and results for code_interpreter_results
|
||||
self._server_tool_inputs: Dict[str, Any] = {}
|
||||
self.tool_results: List[Dict[str, Any]] = []
|
||||
self._current_server_tool_id: Optional[str] = None
|
||||
self._container_id: Optional[str] = None
|
||||
|
||||
def check_empty_tool_call_args(self) -> bool:
|
||||
"""
|
||||
Check if the tool call block so far has been an empty string
|
||||
@@ -682,6 +692,39 @@ class ModelResponseIterator:
|
||||
|
||||
return content_block_start
|
||||
|
||||
def _build_code_interpreter_results(self) -> list:
|
||||
"""Convert accumulated tool_results to OutputCodeInterpreterCall objects.
|
||||
|
||||
Called during streaming to produce provider-neutral code_interpreter_results
|
||||
alongside the raw tool_results, so the Responses API layer doesn't need
|
||||
Anthropic-specific knowledge.
|
||||
|
||||
Returns the full cumulative list each time (not incremental), matching
|
||||
how web_search_results works. stream_chunk_builder uses "last value
|
||||
wins" for list-valued provider_specific_fields keys, so the last
|
||||
emission must contain every result.
|
||||
"""
|
||||
results = []
|
||||
for tr in self.tool_results:
|
||||
if tr.get("type") != "bash_code_execution_tool_result":
|
||||
continue
|
||||
call_id = tr.get("tool_use_id", "")
|
||||
content = tr.get("content", {})
|
||||
log_outputs = build_code_interpreter_log_outputs(content)
|
||||
tool_input = self._server_tool_inputs.get(call_id, {})
|
||||
code = tool_input.get("command", "") if isinstance(tool_input, dict) else ""
|
||||
results.append(
|
||||
OutputCodeInterpreterCall(
|
||||
type="code_interpreter_call",
|
||||
id=call_id,
|
||||
code=code,
|
||||
container_id=self._container_id,
|
||||
status="completed",
|
||||
outputs=log_outputs,
|
||||
)
|
||||
)
|
||||
return results
|
||||
|
||||
def chunk_parser(self, chunk: dict) -> ModelResponseStream: # noqa: PLR0915
|
||||
try:
|
||||
type_chunk = chunk.get("type", "") or ""
|
||||
@@ -748,6 +791,23 @@ class ModelResponseIterator:
|
||||
),
|
||||
index=self.tool_index,
|
||||
)
|
||||
# Track server tool use inputs for code_interpreter_results.
|
||||
# The initial input in content_block_start is typically {}
|
||||
# for streaming; the full input arrives via input_json_delta
|
||||
# and is assembled at content_block_stop.
|
||||
if (
|
||||
content_block_start["content_block"]["type"]
|
||||
== "server_tool_use"
|
||||
):
|
||||
self._current_server_tool_id = content_block_start[
|
||||
"content_block"
|
||||
]["id"]
|
||||
tool_input = content_block_start["content_block"].get(
|
||||
"input", {}
|
||||
)
|
||||
self._server_tool_inputs[
|
||||
self._current_server_tool_id
|
||||
] = tool_input
|
||||
# Include caller information if present (for programmatic tool calling)
|
||||
if "caller" in content_block_start["content_block"]:
|
||||
caller_data = content_block_start["content_block"]["caller"]
|
||||
@@ -808,10 +868,12 @@ class ModelResponseIterator:
|
||||
elif content_type != "tool_search_tool_result":
|
||||
# Handle other tool results (code execution, etc.)
|
||||
# Skip tool_search_tool_result as it's internal metadata
|
||||
if not hasattr(self, "tool_results"):
|
||||
self.tool_results = []
|
||||
self.tool_results.append(content_block_start["content_block"])
|
||||
provider_specific_fields["tool_results"] = self.tool_results
|
||||
# Convert to provider-neutral code_interpreter_results
|
||||
provider_specific_fields[
|
||||
"code_interpreter_results"
|
||||
] = self._build_code_interpreter_results()
|
||||
|
||||
elif type_chunk == "content_block_stop":
|
||||
ContentBlockStop(**chunk) # type: ignore
|
||||
@@ -828,6 +890,26 @@ class ModelResponseIterator:
|
||||
),
|
||||
index=self.tool_index,
|
||||
)
|
||||
# Update server_tool_inputs with fully assembled input
|
||||
# from input_json_delta chunks (content_block_start has {})
|
||||
if (
|
||||
self.current_content_block_type == "server_tool_use"
|
||||
and self._current_server_tool_id
|
||||
):
|
||||
args = ""
|
||||
for block in self.content_blocks:
|
||||
if block["delta"]["type"] == "input_json_delta":
|
||||
partial_json = block["delta"].get("partial_json")
|
||||
if isinstance(partial_json, str):
|
||||
args += partial_json
|
||||
if args:
|
||||
try:
|
||||
self._server_tool_inputs[
|
||||
self._current_server_tool_id
|
||||
] = json.loads(args)
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
pass
|
||||
self._current_server_tool_id = None
|
||||
# Reset response_format tool tracking when block stops
|
||||
self.is_response_format_tool = False
|
||||
# Reset current content block type
|
||||
@@ -840,6 +922,17 @@ class ModelResponseIterator:
|
||||
finish_reason, usage, container = self._handle_message_delta(chunk)
|
||||
if container:
|
||||
provider_specific_fields["container"] = container
|
||||
# Store container_id and re-emit code_interpreter_results
|
||||
# so stream_chunk_builder's last-value-wins picks up the
|
||||
# version with container_id populated.
|
||||
container_id = (
|
||||
container.get("id") if isinstance(container, dict) else None
|
||||
)
|
||||
if container_id and self.tool_results:
|
||||
self._container_id = container_id
|
||||
provider_specific_fields[
|
||||
"code_interpreter_results"
|
||||
] = self._build_code_interpreter_results()
|
||||
elif type_chunk == "message_start":
|
||||
"""
|
||||
Anthropic
|
||||
|
||||
@@ -50,6 +50,10 @@ from litellm.types.llms.openai import (
|
||||
OpenAIMcpServerTool,
|
||||
OpenAIWebSearchOptions,
|
||||
)
|
||||
from litellm.types.responses.main import (
|
||||
OutputCodeInterpreterCall,
|
||||
build_code_interpreter_log_outputs,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
CacheCreationTokenDetails,
|
||||
CompletionTokensDetailsWrapper,
|
||||
@@ -1522,7 +1526,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
||||
tool_results = []
|
||||
tool_results.append(content)
|
||||
|
||||
elif content.get("thinking", None) is not None:
|
||||
elif content.get("type") == "thinking":
|
||||
if thinking_blocks is None:
|
||||
thinking_blocks = []
|
||||
thinking_blocks.append(cast(ChatCompletionThinkingBlock, content))
|
||||
@@ -1682,6 +1686,96 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
||||
)
|
||||
return usage
|
||||
|
||||
def _build_code_by_id_map(
|
||||
self, tool_calls: List[ChatCompletionToolCallChunk]
|
||||
) -> Dict[str, str]:
|
||||
code_by_id: Dict[str, str] = {}
|
||||
for tc in tool_calls:
|
||||
try:
|
||||
args = json.loads(tc.get("function", {}).get("arguments", "{}"))
|
||||
call_id = tc.get("id")
|
||||
command = args.get("command", "")
|
||||
if isinstance(call_id, str):
|
||||
code_by_id[call_id] = command if isinstance(command, str) else ""
|
||||
except Exception:
|
||||
pass
|
||||
return code_by_id
|
||||
|
||||
def _build_code_interpreter_results(
|
||||
self,
|
||||
tool_results: List[Any],
|
||||
code_by_id: Dict[str, str],
|
||||
container_id: Optional[str],
|
||||
) -> List[OutputCodeInterpreterCall]:
|
||||
code_interpreter_results = []
|
||||
for tr in tool_results:
|
||||
if tr.get("type") != "bash_code_execution_tool_result":
|
||||
continue
|
||||
call_id = tr.get("tool_use_id", "")
|
||||
content = tr.get("content", {})
|
||||
log_outputs = build_code_interpreter_log_outputs(content)
|
||||
code_interpreter_results.append(
|
||||
OutputCodeInterpreterCall(
|
||||
type="code_interpreter_call",
|
||||
id=call_id,
|
||||
code=code_by_id.get(call_id, ""),
|
||||
container_id=container_id,
|
||||
status="completed",
|
||||
outputs=log_outputs,
|
||||
)
|
||||
)
|
||||
return code_interpreter_results
|
||||
|
||||
def _build_provider_specific_fields(
|
||||
self,
|
||||
completion_response: dict,
|
||||
citations: Optional[List[Any]],
|
||||
thinking_blocks: Optional[
|
||||
List[
|
||||
Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]
|
||||
]
|
||||
],
|
||||
web_search_results: Optional[List[Any]],
|
||||
tool_results: Optional[List[Any]],
|
||||
compaction_blocks: Optional[List[Any]],
|
||||
tool_calls: List[ChatCompletionToolCallChunk],
|
||||
) -> Dict[str, Any]:
|
||||
provider_specific_fields: Dict[str, Any] = {
|
||||
"citations": citations,
|
||||
"thinking_blocks": thinking_blocks,
|
||||
}
|
||||
|
||||
context_management = completion_response.get("context_management")
|
||||
if context_management is not None:
|
||||
provider_specific_fields["context_management"] = context_management
|
||||
|
||||
if web_search_results is not None:
|
||||
provider_specific_fields["web_search_results"] = web_search_results
|
||||
|
||||
if tool_results is not None:
|
||||
provider_specific_fields["tool_results"] = tool_results
|
||||
container_id = (
|
||||
completion_response.get("container", {}).get("id")
|
||||
if isinstance(completion_response.get("container"), dict)
|
||||
else None
|
||||
)
|
||||
code_by_id = self._build_code_by_id_map(tool_calls)
|
||||
code_interpreter_results = self._build_code_interpreter_results(
|
||||
tool_results, code_by_id, container_id
|
||||
)
|
||||
provider_specific_fields[
|
||||
"code_interpreter_results"
|
||||
] = code_interpreter_results
|
||||
|
||||
container = completion_response.get("container")
|
||||
if container is not None:
|
||||
provider_specific_fields["container"] = container
|
||||
|
||||
if compaction_blocks is not None:
|
||||
provider_specific_fields["compaction_blocks"] = compaction_blocks
|
||||
|
||||
return provider_specific_fields
|
||||
|
||||
def transform_parsed_response(
|
||||
self,
|
||||
completion_response: dict,
|
||||
@@ -1702,98 +1796,73 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
||||
status_code=raw_response.status_code,
|
||||
headers=response_headers,
|
||||
)
|
||||
else:
|
||||
text_content = ""
|
||||
citations: Optional[List[Any]] = None
|
||||
thinking_blocks: Optional[
|
||||
List[
|
||||
Union[
|
||||
ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock
|
||||
]
|
||||
]
|
||||
] = None
|
||||
reasoning_content: Optional[str] = None
|
||||
tool_calls: List[ChatCompletionToolCallChunk] = []
|
||||
|
||||
(
|
||||
text_content,
|
||||
citations,
|
||||
thinking_blocks,
|
||||
reasoning_content,
|
||||
tool_calls,
|
||||
web_search_results,
|
||||
tool_results,
|
||||
compaction_blocks,
|
||||
) = self.extract_response_content(completion_response=completion_response)
|
||||
(
|
||||
text_content,
|
||||
citations,
|
||||
thinking_blocks,
|
||||
reasoning_content,
|
||||
tool_calls,
|
||||
web_search_results,
|
||||
tool_results,
|
||||
compaction_blocks,
|
||||
) = self.extract_response_content(completion_response=completion_response)
|
||||
|
||||
if (
|
||||
prefix_prompt is not None
|
||||
and not text_content.startswith(prefix_prompt)
|
||||
and not litellm.disable_add_prefix_to_prompt
|
||||
):
|
||||
text_content = prefix_prompt + text_content
|
||||
if (
|
||||
prefix_prompt is not None
|
||||
and not text_content.startswith(prefix_prompt)
|
||||
and not litellm.disable_add_prefix_to_prompt
|
||||
):
|
||||
text_content = prefix_prompt + text_content
|
||||
|
||||
context_management: Optional[Dict] = completion_response.get(
|
||||
"context_management"
|
||||
)
|
||||
provider_specific_fields = self._build_provider_specific_fields(
|
||||
completion_response,
|
||||
citations,
|
||||
thinking_blocks,
|
||||
web_search_results,
|
||||
tool_results,
|
||||
compaction_blocks,
|
||||
tool_calls,
|
||||
)
|
||||
|
||||
container: Optional[Dict] = completion_response.get("container")
|
||||
_message = litellm.Message(
|
||||
tool_calls=tool_calls,
|
||||
content=text_content or None,
|
||||
provider_specific_fields=provider_specific_fields,
|
||||
thinking_blocks=thinking_blocks,
|
||||
reasoning_content=reasoning_content,
|
||||
)
|
||||
_message.provider_specific_fields = provider_specific_fields
|
||||
|
||||
provider_specific_fields: Dict[str, Any] = {
|
||||
"citations": citations,
|
||||
"thinking_blocks": thinking_blocks,
|
||||
}
|
||||
if context_management is not None:
|
||||
provider_specific_fields["context_management"] = context_management
|
||||
if web_search_results is not None:
|
||||
provider_specific_fields["web_search_results"] = web_search_results
|
||||
if tool_results is not None:
|
||||
provider_specific_fields["tool_results"] = tool_results
|
||||
if container is not None:
|
||||
provider_specific_fields["container"] = container
|
||||
if compaction_blocks is not None:
|
||||
provider_specific_fields["compaction_blocks"] = compaction_blocks
|
||||
json_mode_message = self._transform_response_for_json_mode(
|
||||
json_mode=json_mode,
|
||||
tool_calls=tool_calls,
|
||||
)
|
||||
if json_mode_message is not None:
|
||||
completion_response["stop_reason"] = "stop"
|
||||
_message = json_mode_message
|
||||
|
||||
_message = litellm.Message(
|
||||
tool_calls=tool_calls,
|
||||
content=text_content or None,
|
||||
provider_specific_fields=provider_specific_fields,
|
||||
thinking_blocks=thinking_blocks,
|
||||
reasoning_content=reasoning_content,
|
||||
)
|
||||
_message.provider_specific_fields = provider_specific_fields
|
||||
model_response.choices[0].message = _message
|
||||
model_response._hidden_params["original_response"] = completion_response[
|
||||
"content"
|
||||
]
|
||||
model_response.choices[0].finish_reason = cast(
|
||||
OpenAIChatCompletionFinishReason,
|
||||
map_finish_reason(completion_response["stop_reason"]),
|
||||
)
|
||||
|
||||
## HANDLE JSON MODE - anthropic returns single function call
|
||||
json_mode_message = self._transform_response_for_json_mode(
|
||||
json_mode=json_mode,
|
||||
tool_calls=tool_calls,
|
||||
)
|
||||
if json_mode_message is not None:
|
||||
completion_response["stop_reason"] = "stop"
|
||||
_message = json_mode_message
|
||||
|
||||
model_response.choices[0].message = _message # type: ignore
|
||||
model_response._hidden_params["original_response"] = completion_response[
|
||||
"content"
|
||||
] # allow user to access raw anthropic tool calling response
|
||||
|
||||
model_response.choices[0].finish_reason = cast(
|
||||
OpenAIChatCompletionFinishReason,
|
||||
map_finish_reason(completion_response["stop_reason"]),
|
||||
)
|
||||
|
||||
## CALCULATING USAGE
|
||||
usage = self.calculate_usage(
|
||||
usage_object=completion_response["usage"],
|
||||
reasoning_content=reasoning_content,
|
||||
completion_response=completion_response,
|
||||
speed=speed,
|
||||
)
|
||||
setattr(model_response, "usage", usage) # type: ignore
|
||||
setattr(model_response, "usage", usage)
|
||||
|
||||
model_response.created = int(time.time())
|
||||
model_response.model = completion_response["model"]
|
||||
|
||||
_hidden_params["provider_specific_fields"] = provider_specific_fields
|
||||
model_response._hidden_params = _hidden_params
|
||||
return model_response
|
||||
|
||||
|
||||
@@ -4462,6 +4462,78 @@
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"azure/gpt-5.4-mini": {
|
||||
"cache_read_input_token_cost": 7.5e-08,
|
||||
"input_cost_per_token": 7.5e-07,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 4.5e-06,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_service_tier": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_none_reasoning_effort": false,
|
||||
"supports_xhigh_reasoning_effort": false
|
||||
},
|
||||
"azure/gpt-5.4-nano": {
|
||||
"cache_read_input_token_cost": 2e-08,
|
||||
"input_cost_per_token": 2e-07,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.25e-06,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_service_tier": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_none_reasoning_effort": false,
|
||||
"supports_xhigh_reasoning_effort": false
|
||||
},
|
||||
"azure/gpt-image-1": {
|
||||
"cache_read_input_image_token_cost": 2.5e-06,
|
||||
"cache_read_input_token_cost": 1.25e-06,
|
||||
|
||||
@@ -208,26 +208,52 @@ async def exchange_token_with_server(
|
||||
client_id: str,
|
||||
client_secret: Optional[str],
|
||||
code_verifier: Optional[str],
|
||||
refresh_token: Optional[str] = None,
|
||||
scope: Optional[str] = None,
|
||||
):
|
||||
if grant_type != "authorization_code":
|
||||
if grant_type not in ("authorization_code", "refresh_token"):
|
||||
raise HTTPException(status_code=400, detail="Unsupported grant_type")
|
||||
|
||||
if mcp_server.token_url is None:
|
||||
raise HTTPException(status_code=400, detail="MCP server token url is not set")
|
||||
|
||||
proxy_base_url = get_request_base_url(request)
|
||||
token_data = {
|
||||
"grant_type": "authorization_code",
|
||||
"client_id": mcp_server.client_id if mcp_server.client_id else client_id,
|
||||
"client_secret": mcp_server.client_secret
|
||||
if mcp_server.client_secret
|
||||
else client_secret,
|
||||
"code": code,
|
||||
"redirect_uri": f"{proxy_base_url}/callback",
|
||||
}
|
||||
resolved_client_id = mcp_server.client_id if mcp_server.client_id else client_id
|
||||
resolved_client_secret = (
|
||||
mcp_server.client_secret if mcp_server.client_secret else client_secret
|
||||
)
|
||||
|
||||
if code_verifier:
|
||||
token_data["code_verifier"] = code_verifier
|
||||
if grant_type == "refresh_token":
|
||||
if not refresh_token:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="refresh_token is required for refresh_token grant",
|
||||
)
|
||||
token_data: dict = {
|
||||
"grant_type": "refresh_token",
|
||||
"refresh_token": refresh_token,
|
||||
"client_id": resolved_client_id,
|
||||
}
|
||||
if resolved_client_secret is not None:
|
||||
token_data["client_secret"] = resolved_client_secret
|
||||
if scope:
|
||||
token_data["scope"] = scope
|
||||
else:
|
||||
if not code:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="code is required for authorization_code grant",
|
||||
)
|
||||
proxy_base_url = get_request_base_url(request)
|
||||
token_data = {
|
||||
"grant_type": "authorization_code",
|
||||
"client_id": resolved_client_id,
|
||||
"code": code,
|
||||
"redirect_uri": f"{proxy_base_url}/callback",
|
||||
}
|
||||
if resolved_client_secret is not None:
|
||||
token_data["client_secret"] = resolved_client_secret
|
||||
if code_verifier:
|
||||
token_data["code_verifier"] = code_verifier
|
||||
|
||||
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
|
||||
response = await async_client.post(
|
||||
@@ -375,6 +401,8 @@ async def token_endpoint(
|
||||
client_id: str = Form(...),
|
||||
client_secret: Optional[str] = Form(None),
|
||||
code_verifier: str = Form(None),
|
||||
refresh_token: Optional[str] = Form(None),
|
||||
scope: Optional[str] = Form(None),
|
||||
mcp_server_name: Optional[str] = None,
|
||||
):
|
||||
"""
|
||||
@@ -408,6 +436,8 @@ async def token_endpoint(
|
||||
client_id=client_id,
|
||||
client_secret=client_secret,
|
||||
code_verifier=code_verifier,
|
||||
refresh_token=refresh_token,
|
||||
scope=scope,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -2955,7 +2955,9 @@ class LiteLLM_ErrorLogs(LiteLLMPydanticObjectBase):
|
||||
endTime: Union[str, datetime, None]
|
||||
|
||||
|
||||
AUDIT_ACTIONS = Literal["created", "updated", "deleted", "blocked", "rotated"]
|
||||
AUDIT_ACTIONS = Literal[
|
||||
"created", "updated", "deleted", "blocked", "unblocked", "rotated"
|
||||
]
|
||||
|
||||
|
||||
class LiteLLM_AuditLogs(LiteLLMPydanticObjectBase):
|
||||
|
||||
@@ -29,6 +29,7 @@ from litellm.constants import (
|
||||
DEFAULT_MAX_RECURSE_DEPTH,
|
||||
EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE,
|
||||
)
|
||||
from litellm.litellm_core_utils.dd_tracing import tracer
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
from litellm.proxy._types import (
|
||||
RBAC_ROLES,
|
||||
@@ -407,18 +408,21 @@ async def common_checks( # noqa: PLR0915
|
||||
|
||||
# 2. If team can call model
|
||||
if _model and team_object:
|
||||
if not await can_team_access_model(
|
||||
model=_model,
|
||||
team_object=team_object,
|
||||
llm_router=llm_router,
|
||||
team_model_aliases=valid_token.team_model_aliases if valid_token else None,
|
||||
):
|
||||
raise ProxyException(
|
||||
message=f"Team not allowed to access model. Team={team_object.team_id}, Model={_model}. Allowed team models = {team_object.models}",
|
||||
type=ProxyErrorTypes.team_model_access_denied,
|
||||
param="model",
|
||||
code=status.HTTP_401_UNAUTHORIZED,
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.common_checks.can_team_access_model"):
|
||||
if not await can_team_access_model(
|
||||
model=_model,
|
||||
team_object=team_object,
|
||||
llm_router=llm_router,
|
||||
team_model_aliases=valid_token.team_model_aliases
|
||||
if valid_token
|
||||
else None,
|
||||
):
|
||||
raise ProxyException(
|
||||
message=f"Team not allowed to access model. Team={team_object.team_id}, Model={_model}. Allowed team models = {team_object.models}",
|
||||
type=ProxyErrorTypes.team_model_access_denied,
|
||||
param="model",
|
||||
code=status.HTTP_401_UNAUTHORIZED,
|
||||
)
|
||||
|
||||
# Require trace id for agent keys when agent has require_trace_id_on_calls_by_agent
|
||||
if valid_token is not None and valid_token.agent_id:
|
||||
@@ -443,54 +447,62 @@ async def common_checks( # noqa: PLR0915
|
||||
|
||||
## 2.1 If user can call model (if personal key)
|
||||
if _model and team_object is None and user_object is not None:
|
||||
await can_user_call_model(
|
||||
model=_model,
|
||||
llm_router=llm_router,
|
||||
user_object=user_object,
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.common_checks.can_user_call_model"):
|
||||
await can_user_call_model(
|
||||
model=_model,
|
||||
llm_router=llm_router,
|
||||
user_object=user_object,
|
||||
)
|
||||
|
||||
# 1.1 - 2.2 - 3.0.2 - 3.0.3: Project checks (blocked, model access, budget)
|
||||
await _run_project_checks(
|
||||
project_object=project_object,
|
||||
_model=_model,
|
||||
llm_router=llm_router,
|
||||
skip_budget_checks=skip_budget_checks,
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.common_checks.run_project_checks"):
|
||||
await _run_project_checks(
|
||||
project_object=project_object,
|
||||
_model=_model,
|
||||
llm_router=llm_router,
|
||||
skip_budget_checks=skip_budget_checks,
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
# If this is a free model, skip all budget checks
|
||||
if not skip_budget_checks:
|
||||
# 3. If team is in budget
|
||||
await _team_max_budget_check(
|
||||
team_object=team_object,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
valid_token=valid_token,
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.common_checks.team_max_budget_check"):
|
||||
await _team_max_budget_check(
|
||||
team_object=team_object,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
valid_token=valid_token,
|
||||
)
|
||||
|
||||
# 3.0.5. If team is over soft budget (alert only, doesn't block)
|
||||
await _team_soft_budget_check(
|
||||
team_object=team_object,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
valid_token=valid_token,
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.common_checks.team_soft_budget_check"):
|
||||
await _team_soft_budget_check(
|
||||
team_object=team_object,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
valid_token=valid_token,
|
||||
)
|
||||
|
||||
# 3.1. If organization is in budget
|
||||
await _organization_max_budget_check(
|
||||
valid_token=valid_token,
|
||||
team_object=team_object,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
with tracer.trace(
|
||||
"litellm.proxy.auth.common_checks.organization_max_budget_check"
|
||||
):
|
||||
await _organization_max_budget_check(
|
||||
valid_token=valid_token,
|
||||
team_object=team_object,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
await _tag_max_budget_check(
|
||||
request_body=request_body,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
valid_token=valid_token,
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.common_checks.tag_max_budget_check"):
|
||||
await _tag_max_budget_check(
|
||||
request_body=request_body,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
valid_token=valid_token,
|
||||
)
|
||||
|
||||
# 4. If user is in budget
|
||||
## 4.1 check personal budget, if personal key
|
||||
@@ -508,14 +520,15 @@ async def common_checks( # noqa: PLR0915
|
||||
)
|
||||
|
||||
## 4.2 check team member budget, if team key
|
||||
await _check_team_member_budget(
|
||||
team_object=team_object,
|
||||
user_object=user_object,
|
||||
valid_token=valid_token,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.common_checks.check_team_member_budget"):
|
||||
await _check_team_member_budget(
|
||||
team_object=team_object,
|
||||
user_object=user_object,
|
||||
valid_token=valid_token,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
# 5. If end_user ('user' passed to /chat/completions, /embeddings endpoint) is in budget
|
||||
if (
|
||||
@@ -554,19 +567,21 @@ async def common_checks( # noqa: PLR0915
|
||||
)
|
||||
|
||||
# 11. [OPTIONAL] Vector store checks - is the object allowed to access the vector store
|
||||
await vector_store_access_check(
|
||||
request_body=request_body,
|
||||
team_object=team_object,
|
||||
valid_token=valid_token,
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.common_checks.vector_store_access_check"):
|
||||
await vector_store_access_check(
|
||||
request_body=request_body,
|
||||
team_object=team_object,
|
||||
valid_token=valid_token,
|
||||
)
|
||||
|
||||
# 12. [OPTIONAL] Tool allowlist - key/team allowed_tools (no DB in hot path)
|
||||
await check_tools_allowlist(
|
||||
request_body=request_body,
|
||||
valid_token=valid_token,
|
||||
team_object=team_object,
|
||||
route=route,
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.common_checks.check_tools_allowlist"):
|
||||
await check_tools_allowlist(
|
||||
request_body=request_body,
|
||||
valid_token=valid_token,
|
||||
team_object=team_object,
|
||||
route=route,
|
||||
)
|
||||
|
||||
return True
|
||||
|
||||
|
||||
@@ -548,13 +548,12 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
||||
custom_auth_api_key: bool = False
|
||||
|
||||
try:
|
||||
# get the request body
|
||||
|
||||
await pre_db_read_auth_checks(
|
||||
request_data=request_data,
|
||||
request=request,
|
||||
route=route,
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.pre_db_read_auth_checks"):
|
||||
await pre_db_read_auth_checks(
|
||||
request_data=request_data,
|
||||
request=request,
|
||||
route=route,
|
||||
)
|
||||
pass_through_endpoints: Optional[List[dict]] = general_settings.get(
|
||||
"pass_through_endpoints", None
|
||||
)
|
||||
@@ -588,9 +587,10 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
||||
|
||||
### USER-DEFINED AUTH FUNCTION ###
|
||||
if enterprise_custom_auth is not None:
|
||||
response = await enterprise_custom_auth(
|
||||
request=request, api_key=api_key, user_custom_auth=user_custom_auth
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.enterprise_custom_auth"):
|
||||
response = await enterprise_custom_auth(
|
||||
request=request, api_key=api_key, user_custom_auth=user_custom_auth
|
||||
)
|
||||
if response is not None and isinstance(response, UserAPIKeyAuth):
|
||||
validated = UserAPIKeyAuth.model_validate(response)
|
||||
validated = await _run_post_custom_auth_checks(
|
||||
@@ -706,18 +706,19 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
||||
# Fall through to virtual key checks
|
||||
|
||||
if do_standard_jwt_auth:
|
||||
result = await JWTAuthManager.auth_builder(
|
||||
request_data=request_data,
|
||||
general_settings=general_settings,
|
||||
api_key=api_key,
|
||||
jwt_handler=jwt_handler,
|
||||
route=route,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
parent_otel_span=parent_otel_span,
|
||||
request_headers=_safe_get_request_headers(request),
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.jwt_auth_builder"):
|
||||
result = await JWTAuthManager.auth_builder(
|
||||
request_data=request_data,
|
||||
general_settings=general_settings,
|
||||
api_key=api_key,
|
||||
jwt_handler=jwt_handler,
|
||||
route=route,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
parent_otel_span=parent_otel_span,
|
||||
request_headers=_safe_get_request_headers(request),
|
||||
)
|
||||
|
||||
is_proxy_admin = result["is_proxy_admin"]
|
||||
team_id = result["team_id"]
|
||||
@@ -909,15 +910,15 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
||||
try:
|
||||
end_user_params["end_user_id"] = end_user_id
|
||||
|
||||
# get end-user object
|
||||
_end_user_object = 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,
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.get_end_user_object"):
|
||||
_end_user_object = 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_object is not None:
|
||||
end_user_params[
|
||||
"allowed_model_region"
|
||||
@@ -960,14 +961,15 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
||||
if valid_token is None:
|
||||
## Check CACHE
|
||||
try:
|
||||
valid_token = await get_key_object(
|
||||
hashed_token=hash_token(api_key),
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_cache_only=True,
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.get_key_object_check_cache"):
|
||||
valid_token = await get_key_object(
|
||||
hashed_token=hash_token(api_key),
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
check_cache_only=True,
|
||||
)
|
||||
except Exception:
|
||||
verbose_logger.debug("api key not found in cache.")
|
||||
valid_token = None
|
||||
@@ -1139,13 +1141,14 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
||||
api_key = hash_token(token=api_key)
|
||||
|
||||
try:
|
||||
valid_token = await get_key_object(
|
||||
hashed_token=api_key,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.get_key_object_from_db"):
|
||||
valid_token = await get_key_object(
|
||||
hashed_token=api_key,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except ProxyException as e:
|
||||
if e.code == 401 or e.code == "401":
|
||||
e.message = "Authentication Error, Invalid proxy server token passed. Received API Key = {}, Key Hash (Token) ={}. Unable to find token in cache or `LiteLLM_VerificationTokenTable`".format(
|
||||
@@ -1233,14 +1236,15 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
||||
# Check 2. If user_id for this token is in budget - done in common_checks()
|
||||
if valid_token.user_id is not None:
|
||||
try:
|
||||
user_obj = await get_user_object(
|
||||
user_id=valid_token.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,
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.get_user_object"):
|
||||
user_obj = await get_user_object(
|
||||
user_id=valid_token.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,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
"litellm.proxy.auth.user_api_key_auth.py::user_api_key_auth() - Unable to get user from db/cache. Setting user_obj to None. Exception received - {}".format(
|
||||
@@ -1329,71 +1333,73 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
||||
)
|
||||
|
||||
if not skip_budget_checks:
|
||||
# Check 4. Token Spend is under budget
|
||||
if RouteChecks.is_llm_api_route(route=route):
|
||||
await _virtual_key_max_budget_check(
|
||||
with tracer.trace("litellm.proxy.auth.budget_checks"):
|
||||
# Check 4. Token Spend is under budget
|
||||
if RouteChecks.is_llm_api_route(route=route):
|
||||
await _virtual_key_max_budget_check(
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_obj=user_obj,
|
||||
)
|
||||
|
||||
# Check 5. Max Budget Alert Check
|
||||
await _virtual_key_max_budget_alert_check(
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_obj=user_obj,
|
||||
)
|
||||
|
||||
# Check 5. Max Budget Alert Check
|
||||
await _virtual_key_max_budget_alert_check(
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_obj=user_obj,
|
||||
)
|
||||
|
||||
# Check 6. Soft Budget Check
|
||||
await _virtual_key_soft_budget_check(
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_obj=user_obj,
|
||||
)
|
||||
|
||||
# Check 5. Token Model Spend is under Model budget
|
||||
max_budget_per_model = valid_token.model_max_budget
|
||||
current_model = request_data.get("model", None)
|
||||
|
||||
if (
|
||||
max_budget_per_model is not None
|
||||
and isinstance(max_budget_per_model, dict)
|
||||
and len(max_budget_per_model) > 0
|
||||
and prisma_client is not None
|
||||
and current_model is not None
|
||||
and valid_token.token is not None
|
||||
):
|
||||
## GET THE SPEND FOR THIS MODEL
|
||||
await model_max_budget_limiter.is_key_within_model_budget(
|
||||
user_api_key_dict=valid_token,
|
||||
model=current_model,
|
||||
# Check 6. Soft Budget Check
|
||||
await _virtual_key_soft_budget_check(
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_obj=user_obj,
|
||||
)
|
||||
|
||||
# Check 5b. End-user model max budget
|
||||
end_user_mmb = valid_token.end_user_model_max_budget
|
||||
if (
|
||||
end_user_mmb is not None
|
||||
and isinstance(end_user_mmb, dict)
|
||||
and len(end_user_mmb) > 0
|
||||
and current_model is not None
|
||||
and valid_token.end_user_id is not None
|
||||
):
|
||||
await model_max_budget_limiter.is_end_user_within_model_budget(
|
||||
end_user_id=valid_token.end_user_id,
|
||||
end_user_model_max_budget=end_user_mmb,
|
||||
model=current_model,
|
||||
)
|
||||
# Check 5. Token Model Spend is under Model budget
|
||||
max_budget_per_model = valid_token.model_max_budget
|
||||
current_model = request_data.get("model", None)
|
||||
|
||||
if (
|
||||
max_budget_per_model is not None
|
||||
and isinstance(max_budget_per_model, dict)
|
||||
and len(max_budget_per_model) > 0
|
||||
and prisma_client is not None
|
||||
and current_model is not None
|
||||
and valid_token.token is not None
|
||||
):
|
||||
## GET THE SPEND FOR THIS MODEL
|
||||
await model_max_budget_limiter.is_key_within_model_budget(
|
||||
user_api_key_dict=valid_token,
|
||||
model=current_model,
|
||||
)
|
||||
|
||||
# Check 5b. End-user model max budget
|
||||
end_user_mmb = valid_token.end_user_model_max_budget
|
||||
if (
|
||||
end_user_mmb is not None
|
||||
and isinstance(end_user_mmb, dict)
|
||||
and len(end_user_mmb) > 0
|
||||
and current_model is not None
|
||||
and valid_token.end_user_id is not None
|
||||
):
|
||||
await model_max_budget_limiter.is_end_user_within_model_budget(
|
||||
end_user_id=valid_token.end_user_id,
|
||||
end_user_model_max_budget=end_user_mmb,
|
||||
model=current_model,
|
||||
)
|
||||
|
||||
# Check 6: Additional Common Checks across jwt + key auth
|
||||
if valid_token.team_id is not None:
|
||||
try:
|
||||
_team_obj = await get_team_object(
|
||||
team_id=valid_token.team_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,
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.get_team_object"):
|
||||
_team_obj = await get_team_object(
|
||||
team_id=valid_token.team_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,
|
||||
)
|
||||
except HTTPException:
|
||||
_team_obj = LiteLLM_TeamTableCachedObj(
|
||||
team_id=valid_token.team_id,
|
||||
@@ -1431,11 +1437,14 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
||||
litellm.max_budget > 0 and prisma_client is not None
|
||||
): # user set proxy max budget
|
||||
cache_key = "{}:spend".format(litellm_proxy_admin_name)
|
||||
global_proxy_spend = await _fetch_global_spend_with_event_coordination(
|
||||
cache_key=cache_key,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.get_global_proxy_spend"):
|
||||
global_proxy_spend = (
|
||||
await _fetch_global_spend_with_event_coordination(
|
||||
cache_key=cache_key,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
)
|
||||
|
||||
if global_proxy_spend is not None:
|
||||
call_info = CallInfo(
|
||||
@@ -1452,21 +1461,22 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
||||
user_info=call_info,
|
||||
)
|
||||
)
|
||||
_ = await common_checks(
|
||||
request=request,
|
||||
request_body=request_data,
|
||||
team_object=_team_obj,
|
||||
user_object=user_obj,
|
||||
end_user_object=_end_user_object,
|
||||
general_settings=general_settings,
|
||||
global_proxy_spend=global_proxy_spend,
|
||||
route=route,
|
||||
llm_router=llm_router,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
valid_token=valid_token,
|
||||
skip_budget_checks=skip_budget_checks,
|
||||
project_object=_project_obj,
|
||||
)
|
||||
with tracer.trace("litellm.proxy.auth.common_checks"):
|
||||
_ = await common_checks(
|
||||
request=request,
|
||||
request_body=request_data,
|
||||
team_object=_team_obj,
|
||||
user_object=user_obj,
|
||||
end_user_object=_end_user_object,
|
||||
general_settings=general_settings,
|
||||
global_proxy_spend=global_proxy_spend,
|
||||
route=route,
|
||||
llm_router=llm_router,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
valid_token=valid_token,
|
||||
skip_budget_checks=skip_budget_checks,
|
||||
project_object=_project_obj,
|
||||
)
|
||||
# Token passed all checks
|
||||
if valid_token is None:
|
||||
raise HTTPException(401, detail="Invalid API key")
|
||||
|
||||
@@ -1260,7 +1260,9 @@ class ProxyBaseLLMRequestProcessing:
|
||||
custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_id=(
|
||||
_litellm_logging_obj.litellm_call_id if _litellm_logging_obj else None
|
||||
_litellm_logging_obj.litellm_call_id
|
||||
if _litellm_logging_obj
|
||||
else self.data.get("litellm_call_id")
|
||||
),
|
||||
model_id=model_id,
|
||||
version=version,
|
||||
|
||||
@@ -41,10 +41,8 @@ from litellm.proxy._experimental.mcp_server.db import (
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy._types import LiteLLM_VerificationToken
|
||||
from litellm.proxy.auth.auth_checks import (
|
||||
_cache_key_object,
|
||||
_delete_cache_key_object,
|
||||
can_team_access_model,
|
||||
get_key_object,
|
||||
get_org_object,
|
||||
get_project_object,
|
||||
get_team_object,
|
||||
@@ -1656,7 +1654,7 @@ async def _get_and_validate_existing_key(
|
||||
LiteLLM_VerificationToken: The existing key row
|
||||
|
||||
Raises:
|
||||
HTTPException: If key is not found
|
||||
ProxyException: 404 if key is not found
|
||||
"""
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
@@ -1664,16 +1662,18 @@ async def _get_and_validate_existing_key(
|
||||
detail={"error": "Database not connected"},
|
||||
)
|
||||
|
||||
existing_key_row = await prisma_client.get_data(
|
||||
token=token,
|
||||
table_name="key",
|
||||
query_type="find_unique",
|
||||
hashed_token = _hash_token_if_needed(token=token)
|
||||
|
||||
existing_key_row = await prisma_client.db.litellm_verificationtoken.find_unique(
|
||||
where={"token": hashed_token}
|
||||
)
|
||||
|
||||
if existing_key_row is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": f"Key not found: {token}"},
|
||||
raise ProxyException(
|
||||
message="Key not found.",
|
||||
type=ProxyErrorTypes.not_found_error,
|
||||
param="key",
|
||||
code=status.HTTP_404_NOT_FOUND,
|
||||
)
|
||||
|
||||
return existing_key_row
|
||||
@@ -2111,19 +2111,11 @@ async def update_key_fn(
|
||||
key = data_json.pop("key")
|
||||
|
||||
# get the row from db
|
||||
if prisma_client is None:
|
||||
raise Exception("Not connected to DB!")
|
||||
|
||||
existing_key_row = await prisma_client.get_data(
|
||||
token=data.key, table_name="key", query_type="find_unique"
|
||||
existing_key_row = await _get_and_validate_existing_key(
|
||||
token=data.key,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
if existing_key_row is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": f"Team not found, passed team_id={data.team_id}"},
|
||||
)
|
||||
|
||||
await _validate_update_key_data(
|
||||
data=data,
|
||||
existing_key_row=existing_key_row,
|
||||
@@ -2158,6 +2150,8 @@ async def update_key_fn(
|
||||
)
|
||||
|
||||
_data = {**non_default_values, "token": key}
|
||||
if prisma_client is None:
|
||||
raise Exception("Not connected to DB!")
|
||||
response = await prisma_client.update_data(token=key, data=_data)
|
||||
|
||||
# Delete - key from cache, since it's been updated!
|
||||
@@ -2330,6 +2324,8 @@ async def bulk_update_keys(
|
||||
error_message = error_detail.get("error", str(e))
|
||||
else:
|
||||
error_message = str(error_detail)
|
||||
elif isinstance(e, ProxyException):
|
||||
error_message = e.message
|
||||
else:
|
||||
error_message = str(e)
|
||||
|
||||
@@ -4945,18 +4941,19 @@ async def block_key(
|
||||
route="/key/block",
|
||||
)
|
||||
|
||||
if litellm.store_audit_logs is True:
|
||||
# make an audit log for key update
|
||||
record = await prisma_client.db.litellm_verificationtoken.find_unique(
|
||||
where={"token": hashed_token}
|
||||
# Check if the key exists before trying to block it
|
||||
existing_record = await prisma_client.db.litellm_verificationtoken.find_unique(
|
||||
where={"token": hashed_token}
|
||||
)
|
||||
if existing_record is None:
|
||||
raise ProxyException(
|
||||
message="Key not found.",
|
||||
type=ProxyErrorTypes.not_found_error,
|
||||
param="key",
|
||||
code=status.HTTP_404_NOT_FOUND,
|
||||
)
|
||||
if record is None:
|
||||
raise ProxyException(
|
||||
message=f"Key {data.key} not found",
|
||||
type=ProxyErrorTypes.bad_request_error,
|
||||
param="key",
|
||||
code=status.HTTP_404_NOT_FOUND,
|
||||
)
|
||||
|
||||
if litellm.store_audit_logs is True:
|
||||
asyncio.create_task(
|
||||
create_audit_log_for_update(
|
||||
request_data=LiteLLM_AuditLogs(
|
||||
@@ -4970,7 +4967,7 @@ async def block_key(
|
||||
object_id=hashed_token,
|
||||
action="blocked",
|
||||
updated_values="{}",
|
||||
before_value=record.model_dump_json(),
|
||||
before_value=existing_record.model_dump_json(),
|
||||
)
|
||||
)
|
||||
)
|
||||
@@ -4979,24 +4976,9 @@ async def block_key(
|
||||
where={"token": hashed_token}, data={"blocked": True} # type: ignore
|
||||
)
|
||||
|
||||
## UPDATE KEY CACHE
|
||||
|
||||
### get cached object ###
|
||||
key_object = await get_key_object(
|
||||
## UPDATE KEY CACHE - invalidate so next read re-fetches from DB
|
||||
await _delete_cache_key_object(
|
||||
hashed_token=hashed_token,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
### update cached object ###
|
||||
key_object.blocked = True
|
||||
|
||||
### store cached object ###
|
||||
await _cache_key_object(
|
||||
hashed_token=hashed_token,
|
||||
user_api_key_obj=key_object,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
@@ -5068,18 +5050,19 @@ async def unblock_key(
|
||||
route="/key/unblock",
|
||||
)
|
||||
|
||||
if litellm.store_audit_logs is True:
|
||||
# make an audit log for key update
|
||||
record = await prisma_client.db.litellm_verificationtoken.find_unique(
|
||||
where={"token": hashed_token}
|
||||
# Check if the key exists before trying to unblock it
|
||||
existing_record = await prisma_client.db.litellm_verificationtoken.find_unique(
|
||||
where={"token": hashed_token}
|
||||
)
|
||||
if existing_record is None:
|
||||
raise ProxyException(
|
||||
message="Key not found.",
|
||||
type=ProxyErrorTypes.not_found_error,
|
||||
param="key",
|
||||
code=status.HTTP_404_NOT_FOUND,
|
||||
)
|
||||
if record is None:
|
||||
raise ProxyException(
|
||||
message=f"Key {data.key} not found",
|
||||
type=ProxyErrorTypes.bad_request_error,
|
||||
param="key",
|
||||
code=status.HTTP_404_NOT_FOUND,
|
||||
)
|
||||
|
||||
if litellm.store_audit_logs is True:
|
||||
asyncio.create_task(
|
||||
create_audit_log_for_update(
|
||||
request_data=LiteLLM_AuditLogs(
|
||||
@@ -5091,9 +5074,9 @@ async def unblock_key(
|
||||
changed_by_api_key=user_api_key_dict.api_key,
|
||||
table_name=LitellmTableNames.KEY_TABLE_NAME,
|
||||
object_id=hashed_token,
|
||||
action="blocked",
|
||||
action="unblocked",
|
||||
updated_values="{}",
|
||||
before_value=record.model_dump_json(),
|
||||
before_value=existing_record.model_dump_json(),
|
||||
)
|
||||
)
|
||||
)
|
||||
@@ -5102,24 +5085,9 @@ async def unblock_key(
|
||||
where={"token": hashed_token}, data={"blocked": False} # type: ignore
|
||||
)
|
||||
|
||||
## UPDATE KEY CACHE
|
||||
|
||||
### get cached object ###
|
||||
key_object = await get_key_object(
|
||||
## UPDATE KEY CACHE - invalidate so next read re-fetches from DB
|
||||
await _delete_cache_key_object(
|
||||
hashed_token=hashed_token,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=None,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
### update cached object ###
|
||||
key_object.blocked = False
|
||||
|
||||
### store cached object ###
|
||||
await _cache_key_object(
|
||||
hashed_token=hashed_token,
|
||||
user_api_key_obj=key_object,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
@@ -1399,6 +1399,8 @@ if MCP_AVAILABLE:
|
||||
client_id: Optional[str] = Form(None),
|
||||
client_secret: Optional[str] = Form(None),
|
||||
code_verifier: Optional[str] = Form(None),
|
||||
refresh_token: Optional[str] = Form(None),
|
||||
scope: Optional[str] = Form(None),
|
||||
):
|
||||
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id)
|
||||
resolved_client_id = mcp_server.client_id or client_id or ""
|
||||
@@ -1422,6 +1424,8 @@ if MCP_AVAILABLE:
|
||||
client_id=resolved_client_id,
|
||||
client_secret=client_secret,
|
||||
code_verifier=code_verifier,
|
||||
refresh_token=refresh_token,
|
||||
scope=scope,
|
||||
)
|
||||
|
||||
@router.post(
|
||||
|
||||
@@ -3334,7 +3334,9 @@ def _convert_teams_to_response_models(
|
||||
use_deleted_table: bool,
|
||||
) -> List[Union[TeamListItem, LiteLLM_TeamTable, LiteLLM_DeletedTeamTable]]:
|
||||
"""Convert raw Prisma team rows to response models."""
|
||||
team_list: List[Union[TeamListItem, LiteLLM_TeamTable, LiteLLM_DeletedTeamTable]] = []
|
||||
team_list: List[
|
||||
Union[TeamListItem, LiteLLM_TeamTable, LiteLLM_DeletedTeamTable]
|
||||
] = []
|
||||
for team in teams:
|
||||
try:
|
||||
team_dict = team.model_dump()
|
||||
|
||||
+13
-3
@@ -7,6 +7,7 @@ import httpx
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.litellm_logging import use_custom_pricing_for_model
|
||||
from litellm.llms.anthropic import get_anthropic_config
|
||||
from litellm.llms.anthropic.chat.handler import (
|
||||
ModelResponseIterator as AnthropicModelResponseIterator,
|
||||
@@ -124,10 +125,21 @@ class AnthropicPassthroughLoggingHandler:
|
||||
if custom_llm_provider and not model.startswith(f"{custom_llm_provider}/"):
|
||||
model_for_cost = f"{custom_llm_provider}/{model}"
|
||||
|
||||
router_model_id = logging_obj.get_router_model_id()
|
||||
custom_pricing = use_custom_pricing_for_model(
|
||||
litellm_params=(
|
||||
logging_obj.litellm_params
|
||||
if hasattr(logging_obj, "litellm_params")
|
||||
else None
|
||||
)
|
||||
)
|
||||
|
||||
response_cost = litellm.completion_cost(
|
||||
completion_response=litellm_model_response,
|
||||
model=model_for_cost,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
custom_pricing=custom_pricing,
|
||||
router_model_id=router_model_id,
|
||||
)
|
||||
|
||||
kwargs["response_cost"] = response_cost
|
||||
@@ -319,9 +331,7 @@ class AnthropicPassthroughLoggingHandler:
|
||||
import base64
|
||||
|
||||
from litellm._uuid import uuid
|
||||
from litellm.llms.anthropic.batches.transformation import (
|
||||
AnthropicBatchesConfig,
|
||||
)
|
||||
from litellm.llms.anthropic.batches.transformation import AnthropicBatchesConfig
|
||||
from litellm.types.utils import Choices, SpecialEnums
|
||||
|
||||
try:
|
||||
|
||||
@@ -12,6 +12,7 @@ import click
|
||||
import httpx
|
||||
from dotenv import load_dotenv
|
||||
|
||||
import litellm
|
||||
from litellm.constants import DEFAULT_NUM_WORKERS_LITELLM_PROXY
|
||||
from litellm.secret_managers.main import get_secret_bool
|
||||
|
||||
@@ -387,7 +388,7 @@ class ProxyInitializationHelpers:
|
||||
@click.option("--api_base", default=None, help="API base URL.")
|
||||
@click.option(
|
||||
"--api_version",
|
||||
default="2024-07-01-preview",
|
||||
default=litellm.AZURE_DEFAULT_API_VERSION,
|
||||
help="For azure - pass in the api version.",
|
||||
)
|
||||
@click.option(
|
||||
|
||||
@@ -1878,28 +1878,33 @@ class ProxyLogging:
|
||||
)
|
||||
|
||||
input: Union[list, str, dict] = ""
|
||||
normalized_call_type: Optional[str] = None
|
||||
if "messages" in request_data and isinstance(
|
||||
request_data["messages"], list
|
||||
):
|
||||
input = request_data["messages"]
|
||||
litellm_logging_obj.model_call_details["messages"] = input
|
||||
if litellm_logging_obj.call_type != CallTypes.pass_through.value:
|
||||
litellm_logging_obj.call_type = CallTypes.acompletion.value
|
||||
normalized_call_type = CallTypes.acompletion.value
|
||||
elif "prompt" in request_data and isinstance(request_data["prompt"], str):
|
||||
input = request_data["prompt"]
|
||||
litellm_logging_obj.model_call_details["prompt"] = input
|
||||
if litellm_logging_obj.call_type != CallTypes.pass_through.value:
|
||||
litellm_logging_obj.call_type = CallTypes.atext_completion.value
|
||||
normalized_call_type = CallTypes.atext_completion.value
|
||||
elif "input" in request_data and isinstance(request_data["input"], list):
|
||||
input = request_data["input"]
|
||||
litellm_logging_obj.model_call_details["input"] = input
|
||||
if litellm_logging_obj.call_type != CallTypes.pass_through.value:
|
||||
litellm_logging_obj.call_type = CallTypes.aembedding.value
|
||||
normalized_call_type = CallTypes.aembedding.value
|
||||
if normalized_call_type is not None:
|
||||
litellm_logging_obj.call_type = normalized_call_type
|
||||
litellm_logging_obj.model_call_details[
|
||||
"call_type"
|
||||
] = normalized_call_type
|
||||
# Pass-through endpoints are logged via the callback loop's
|
||||
# async_post_call_failure_hook — skip pre_call and failure handlers.
|
||||
if litellm_logging_obj.call_type == CallTypes.pass_through.value:
|
||||
return
|
||||
|
||||
litellm_logging_obj.pre_call(
|
||||
input=input,
|
||||
api_key="",
|
||||
|
||||
@@ -107,6 +107,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
||||
self._reasoning_done_emitted = False
|
||||
self._reasoning_item_id: Optional[str] = None
|
||||
self._accumulated_reasoning_content_parts: List[str] = []
|
||||
self._accumulated_provider_specific_fields: Dict[str, Any] = {}
|
||||
|
||||
def _get_or_assign_tool_output_index(self, call_id: str) -> int:
|
||||
existing = self._tool_output_index_by_call_id.get(call_id)
|
||||
@@ -479,16 +480,36 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
||||
event.__dict__["sequence_number"] = self._sequence_number
|
||||
return event
|
||||
|
||||
def create_litellm_model_response(
|
||||
self,
|
||||
) -> Optional[ModelResponse]:
|
||||
return cast(
|
||||
def _merge_provider_specific_fields(self, src: dict) -> None:
|
||||
"""Merge provider_specific_fields using last-value-wins for lists.
|
||||
|
||||
List-valued keys (web_search_results, tool_results,
|
||||
code_interpreter_results, etc.) are emitted cumulatively — each
|
||||
emission contains the full list so far. Using "last value wins"
|
||||
matches stream_chunk_builder's semantics and avoids quadratic
|
||||
growth from repeated extend calls.
|
||||
"""
|
||||
for key, val in src.items():
|
||||
self._accumulated_provider_specific_fields[key] = val
|
||||
|
||||
def create_litellm_model_response(self) -> Optional[ModelResponse]:
|
||||
response = cast(
|
||||
Optional[ModelResponse],
|
||||
stream_chunk_builder(
|
||||
chunks=self.collected_chat_completion_chunks,
|
||||
logging_obj=self.litellm_logging_obj,
|
||||
),
|
||||
)
|
||||
if response is not None and self._accumulated_provider_specific_fields:
|
||||
if (
|
||||
not hasattr(response, "_hidden_params")
|
||||
or response._hidden_params is None
|
||||
):
|
||||
response._hidden_params = {}
|
||||
response._hidden_params.setdefault("provider_specific_fields", {}).update(
|
||||
self._accumulated_provider_specific_fields
|
||||
)
|
||||
return response
|
||||
|
||||
@staticmethod
|
||||
def _snapshot_chunk_for_stream_chunk_builder(
|
||||
@@ -853,6 +874,17 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
||||
if chunk is not None:
|
||||
chunk = cast(ModelResponseStream, chunk)
|
||||
self._ensure_output_item_for_chunk(chunk)
|
||||
# Accumulate provider_specific_fields from chunk and delta
|
||||
for src in (
|
||||
getattr(chunk, "provider_specific_fields", None),
|
||||
getattr(
|
||||
chunk.choices[0].delta if chunk.choices else None,
|
||||
"provider_specific_fields",
|
||||
None,
|
||||
),
|
||||
):
|
||||
if src and isinstance(src, dict):
|
||||
self._merge_provider_specific_fields(src)
|
||||
# Proceed to transformation
|
||||
self.collected_chat_completion_chunks.append(
|
||||
self._snapshot_chunk_for_stream_chunk_builder(chunk)
|
||||
@@ -964,6 +996,17 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
||||
try:
|
||||
chunk = self.litellm_custom_stream_wrapper.__next__()
|
||||
self._ensure_output_item_for_chunk(chunk)
|
||||
# Accumulate provider_specific_fields from chunk and delta
|
||||
for src in (
|
||||
getattr(chunk, "provider_specific_fields", None),
|
||||
getattr(
|
||||
chunk.choices[0].delta if chunk.choices else None,
|
||||
"provider_specific_fields",
|
||||
None,
|
||||
),
|
||||
):
|
||||
if src and isinstance(src, dict):
|
||||
self._merge_provider_specific_fields(src)
|
||||
# Emit any just-queued output_item event
|
||||
if self._pending_response_events:
|
||||
return self._pending_response_events.pop(0)
|
||||
|
||||
@@ -42,6 +42,7 @@ from litellm.types.llms.openai import (
|
||||
from litellm.types.responses.main import (
|
||||
GenericResponseOutputItem,
|
||||
GenericResponseOutputItemContentAnnotation,
|
||||
OutputCodeInterpreterCall,
|
||||
OutputFunctionToolCall,
|
||||
OutputImageGenerationCall,
|
||||
OutputText,
|
||||
@@ -1696,6 +1697,7 @@ class LiteLLMCompletionResponsesConfig:
|
||||
) -> List[
|
||||
Union[
|
||||
GenericResponseOutputItem,
|
||||
OutputCodeInterpreterCall,
|
||||
OutputFunctionToolCall,
|
||||
OutputImageGenerationCall,
|
||||
ResponseFunctionToolCall,
|
||||
@@ -1704,6 +1706,7 @@ class LiteLLMCompletionResponsesConfig:
|
||||
responses_output: List[
|
||||
Union[
|
||||
GenericResponseOutputItem,
|
||||
OutputCodeInterpreterCall,
|
||||
OutputFunctionToolCall,
|
||||
OutputImageGenerationCall,
|
||||
ResponseFunctionToolCall,
|
||||
@@ -1725,8 +1728,63 @@ class LiteLLMCompletionResponsesConfig:
|
||||
chat_completion_response=chat_completion_response
|
||||
)
|
||||
)
|
||||
|
||||
# Convert server-side tool results (e.g. Anthropic code execution)
|
||||
# into code_interpreter_call output items, replacing the corresponding
|
||||
# function_call items so the output matches OpenAI's native shape.
|
||||
tool_result_items = (
|
||||
LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(
|
||||
chat_completion_response
|
||||
)
|
||||
)
|
||||
if tool_result_items:
|
||||
result_by_id = {item.id: item for item in tool_result_items}
|
||||
replaced_ids = set(result_by_id.keys())
|
||||
responses_output = [
|
||||
(
|
||||
result_by_id[getattr(item, "call_id", None)]
|
||||
if (
|
||||
getattr(item, "type", None) == "function_call"
|
||||
and getattr(item, "call_id", None) in replaced_ids
|
||||
)
|
||||
else item
|
||||
)
|
||||
for item in responses_output
|
||||
]
|
||||
|
||||
return responses_output
|
||||
|
||||
@staticmethod
|
||||
def _extract_tool_result_output_items(
|
||||
chat_completion_response: ModelResponse,
|
||||
) -> list:
|
||||
"""Extract pre-built code_interpreter_call output items from provider_specific_fields.
|
||||
|
||||
Provider transformers (e.g. Anthropic) convert their native tool results
|
||||
into OutputCodeInterpreterCall objects and store them in
|
||||
provider_specific_fields["code_interpreter_results"]. This method
|
||||
simply retrieves them — no provider-specific parsing here.
|
||||
"""
|
||||
output_items: list = []
|
||||
for choice in chat_completion_response.choices or []:
|
||||
message = getattr(choice, "message", None)
|
||||
if not message:
|
||||
continue
|
||||
psf = getattr(message, "provider_specific_fields", None)
|
||||
if not psf or not isinstance(psf, dict):
|
||||
continue
|
||||
results = psf.get("code_interpreter_results")
|
||||
if results and isinstance(results, list):
|
||||
for item in results:
|
||||
# In the streaming path, items are plain dicts after
|
||||
# model_dump() in stream_chunk_builder. Reconstruct
|
||||
# Pydantic objects so responses_output has a uniform type.
|
||||
if isinstance(item, dict):
|
||||
output_items.append(OutputCodeInterpreterCall(**item))
|
||||
else:
|
||||
output_items.append(item)
|
||||
return output_items
|
||||
|
||||
@staticmethod
|
||||
def _extract_reasoning_output_items(
|
||||
chat_completion_response: ModelResponse,
|
||||
|
||||
@@ -166,11 +166,12 @@ class BaseResponsesAPIStreamingIterator:
|
||||
)
|
||||
setattr(item, "encrypted_content", wrapped_content)
|
||||
|
||||
# Store the completed response
|
||||
if (
|
||||
openai_responses_api_chunk
|
||||
and getattr(openai_responses_api_chunk, "type", None)
|
||||
== ResponsesAPIStreamEvents.RESPONSE_COMPLETED
|
||||
# Store the completed response (also for incomplete/failed so logging still fires)
|
||||
_chunk_type = getattr(openai_responses_api_chunk, "type", None)
|
||||
if openai_responses_api_chunk and _chunk_type in (
|
||||
ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
|
||||
ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE,
|
||||
ResponsesAPIStreamEvents.RESPONSE_FAILED,
|
||||
):
|
||||
self.completed_response = openai_responses_api_chunk
|
||||
# Add cost to usage object if include_cost_in_streaming_usage is True
|
||||
@@ -195,10 +196,12 @@ class BaseResponsesAPIStreamingIterator:
|
||||
if cost is not None:
|
||||
setattr(usage_obj, "cost", cost)
|
||||
except Exception:
|
||||
# If cost calculation fails, continue without cost
|
||||
pass
|
||||
|
||||
self._handle_logging_completed_response()
|
||||
if _chunk_type == ResponsesAPIStreamEvents.RESPONSE_FAILED:
|
||||
self._handle_logging_failed_response()
|
||||
else:
|
||||
self._handle_logging_completed_response()
|
||||
|
||||
return openai_responses_api_chunk
|
||||
|
||||
@@ -216,6 +219,32 @@ class BaseResponsesAPIStreamingIterator:
|
||||
"""Base implementation - should be overridden by subclasses"""
|
||||
pass
|
||||
|
||||
def _handle_logging_failed_response(self):
|
||||
"""
|
||||
Handle logging for RESPONSE_FAILED events by routing to failure handlers.
|
||||
|
||||
Unlike _handle_logging_completed_response (which calls success handlers),
|
||||
this constructs an exception from the response error and routes to
|
||||
async_failure_handler / failure_handler so logging integrations correctly
|
||||
record the call as failed.
|
||||
"""
|
||||
response_obj = (
|
||||
getattr(self.completed_response, "response", None)
|
||||
if self.completed_response
|
||||
else None
|
||||
)
|
||||
error_info = getattr(response_obj, "error", None) if response_obj else None
|
||||
error_message = "Response failed"
|
||||
if isinstance(error_info, dict):
|
||||
error_message = error_info.get("message", str(error_info))
|
||||
exception = litellm.APIError(
|
||||
status_code=500,
|
||||
message=error_message,
|
||||
llm_provider=self.custom_llm_provider or "",
|
||||
model=self.model or "",
|
||||
)
|
||||
self._handle_failure(exception)
|
||||
|
||||
async def _call_post_streaming_deployment_hook(self, chunk):
|
||||
"""
|
||||
Allow callbacks to modify streaming chunks before returning (parity with chat).
|
||||
|
||||
@@ -3874,14 +3874,23 @@ class Router:
|
||||
The response from the handler function
|
||||
"""
|
||||
handler_name = original_function.__name__
|
||||
metadata_variable_name = _get_router_metadata_variable_name(
|
||||
function_name="generic_api_call"
|
||||
)
|
||||
try:
|
||||
verbose_router_logger.debug(
|
||||
f"Inside _generic_api_call() - handler: {handler_name}, model: {model}; kwargs: {kwargs}"
|
||||
)
|
||||
self._update_kwargs_before_fallbacks(
|
||||
model=model,
|
||||
kwargs=kwargs,
|
||||
metadata_variable_name=metadata_variable_name,
|
||||
)
|
||||
deployment = self.get_available_deployment(
|
||||
model=model,
|
||||
messages=kwargs.get("messages", None),
|
||||
specific_deployment=kwargs.pop("specific_deployment", None),
|
||||
request_kwargs=kwargs,
|
||||
)
|
||||
self._update_kwargs_with_deployment(
|
||||
deployment=deployment, kwargs=kwargs, function_name="generic_api_call"
|
||||
|
||||
@@ -84,6 +84,7 @@ from typing_extensions import Annotated, Dict, Required, TypedDict, override
|
||||
from litellm.types.llms.base import BaseLiteLLMOpenAIResponseObject
|
||||
from litellm.types.responses.main import (
|
||||
GenericResponseOutputItem,
|
||||
OutputCodeInterpreterCall,
|
||||
OutputFunctionToolCall,
|
||||
OutputImageGenerationCall,
|
||||
)
|
||||
@@ -1242,6 +1243,7 @@ class ResponsesAPIResponse(BaseLiteLLMOpenAIResponseObject):
|
||||
List[
|
||||
Union[
|
||||
GenericResponseOutputItem,
|
||||
OutputCodeInterpreterCall,
|
||||
OutputFunctionToolCall,
|
||||
OutputImageGenerationCall,
|
||||
ResponseFunctionToolCall,
|
||||
@@ -1308,13 +1310,16 @@ class ResponsesAPIResponse(BaseLiteLLMOpenAIResponseObject):
|
||||
if not isinstance(serialized, list):
|
||||
return serialized
|
||||
return [
|
||||
{
|
||||
k: v
|
||||
for k, v in item.items()
|
||||
if v is not None or k not in ("status", "content", "encrypted_content")
|
||||
}
|
||||
if isinstance(item, dict) and item.get("type") == "reasoning"
|
||||
else item
|
||||
(
|
||||
{
|
||||
k: v
|
||||
for k, v in item.items()
|
||||
if v is not None
|
||||
or k not in ("status", "content", "encrypted_content")
|
||||
}
|
||||
if isinstance(item, dict) and item.get("type") == "reasoning"
|
||||
else item
|
||||
)
|
||||
for item in serialized
|
||||
]
|
||||
|
||||
|
||||
@@ -49,6 +49,42 @@ class OutputImageGenerationCall(BaseLiteLLMOpenAIResponseObject):
|
||||
result: Optional[str] # Base64 encoded image data (without data:image prefix)
|
||||
|
||||
|
||||
class OutputCodeInterpreterCallLog(BaseLiteLLMOpenAIResponseObject):
|
||||
"""Log output from a code interpreter call"""
|
||||
|
||||
type: Literal["logs"]
|
||||
logs: str
|
||||
|
||||
|
||||
class OutputCodeInterpreterCall(BaseLiteLLMOpenAIResponseObject):
|
||||
"""A code interpreter / code execution call output"""
|
||||
|
||||
type: Literal["code_interpreter_call"]
|
||||
id: str
|
||||
code: Optional[str]
|
||||
container_id: Optional[str]
|
||||
status: Literal["in_progress", "completed", "incomplete", "failed"]
|
||||
outputs: Optional[List[OutputCodeInterpreterCallLog]]
|
||||
|
||||
|
||||
def build_code_interpreter_log_outputs(
|
||||
content: Any,
|
||||
) -> Optional[List[OutputCodeInterpreterCallLog]]:
|
||||
"""Convert Anthropic bash_code_execution stdout/stderr to log outputs.
|
||||
|
||||
Shared by streaming (handler.py) and non-streaming (transformation.py) paths.
|
||||
"""
|
||||
if not isinstance(content, dict):
|
||||
return None
|
||||
parts = []
|
||||
if content.get("stdout"):
|
||||
parts.append(content["stdout"])
|
||||
if content.get("stderr"):
|
||||
parts.append(f"STDERR: {content['stderr']}")
|
||||
logs = "".join(parts)
|
||||
return [OutputCodeInterpreterCallLog(type="logs", logs=logs)] if logs else None
|
||||
|
||||
|
||||
class GenericResponseOutputItem(BaseLiteLLMOpenAIResponseObject):
|
||||
"""
|
||||
Generic response API output item
|
||||
|
||||
@@ -4462,6 +4462,78 @@
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
},
|
||||
"azure/gpt-5.4-mini": {
|
||||
"cache_read_input_token_cost": 7.5e-08,
|
||||
"input_cost_per_token": 7.5e-07,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 4.5e-06,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_service_tier": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_none_reasoning_effort": false,
|
||||
"supports_xhigh_reasoning_effort": false
|
||||
},
|
||||
"azure/gpt-5.4-nano": {
|
||||
"cache_read_input_token_cost": 2e-08,
|
||||
"input_cost_per_token": 2e-07,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.25e-06,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/batch",
|
||||
"/v1/responses"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_pdf_input": true,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_service_tier": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"supports_none_reasoning_effort": false,
|
||||
"supports_xhigh_reasoning_effort": false
|
||||
},
|
||||
"azure/gpt-image-1": {
|
||||
"cache_read_input_image_token_cost": 2.5e-06,
|
||||
"cache_read_input_token_cost": 1.25e-06,
|
||||
@@ -37032,5 +37104,157 @@
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"volcengine/doubao-seed-2-0-pro-260215": {
|
||||
"litellm_provider": "volcengine",
|
||||
"max_input_tokens": 256000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"source": "https://www.volcengine.com/docs/82379/1330310",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_tool_choice": false,
|
||||
"supports_vision": true,
|
||||
"tiered_pricing": [
|
||||
{
|
||||
"input_cost_per_token": 4.6e-07,
|
||||
"output_cost_per_token": 2.3e-06,
|
||||
"range": [
|
||||
0,
|
||||
32000.0
|
||||
]
|
||||
},
|
||||
{
|
||||
"input_cost_per_token": 7e-07,
|
||||
"output_cost_per_token": 3.5e-06,
|
||||
"range": [
|
||||
32000.0,
|
||||
128000.0
|
||||
]
|
||||
},
|
||||
{
|
||||
"input_cost_per_token": 1.4e-06,
|
||||
"output_cost_per_token": 7e-06,
|
||||
"range": [
|
||||
128000.0,
|
||||
256000.0
|
||||
]
|
||||
}
|
||||
]
|
||||
},
|
||||
"volcengine/doubao-seed-2-0-lite-260215": {
|
||||
"litellm_provider": "volcengine",
|
||||
"max_input_tokens": 256000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"source": "https://www.volcengine.com/docs/82379/1330310",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_tool_choice": false,
|
||||
"supports_vision": true,
|
||||
"tiered_pricing": [
|
||||
{
|
||||
"input_cost_per_token": 8.7e-08,
|
||||
"output_cost_per_token": 5.2e-07,
|
||||
"range": [
|
||||
0,
|
||||
32000.0
|
||||
]
|
||||
},
|
||||
{
|
||||
"input_cost_per_token": 1.3e-07,
|
||||
"output_cost_per_token": 7.8e-07,
|
||||
"range": [
|
||||
32000.0,
|
||||
128000.0
|
||||
]
|
||||
},
|
||||
{
|
||||
"input_cost_per_token": 2.6e-07,
|
||||
"output_cost_per_token": 1.6e-06,
|
||||
"range": [
|
||||
128000.0,
|
||||
256000.0
|
||||
]
|
||||
}
|
||||
]
|
||||
},
|
||||
"volcengine/doubao-seed-2-0-mini-260215": {
|
||||
"litellm_provider": "volcengine",
|
||||
"max_input_tokens": 256000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"source": "https://www.volcengine.com/docs/82379/1330310",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_tool_choice": false,
|
||||
"supports_vision": true,
|
||||
"tiered_pricing": [
|
||||
{
|
||||
"input_cost_per_token": 2.9e-08,
|
||||
"output_cost_per_token": 2.9e-07,
|
||||
"range": [
|
||||
0,
|
||||
32000.0
|
||||
]
|
||||
},
|
||||
{
|
||||
"input_cost_per_token": 5.8e-08,
|
||||
"output_cost_per_token": 5.8e-07,
|
||||
"range": [
|
||||
32000.0,
|
||||
128000.0
|
||||
]
|
||||
},
|
||||
{
|
||||
"input_cost_per_token": 1.2e-07,
|
||||
"output_cost_per_token": 1.2e-06,
|
||||
"range": [
|
||||
128000.0,
|
||||
256000.0
|
||||
]
|
||||
}
|
||||
]
|
||||
},
|
||||
"volcengine/doubao-seed-2-0-code-preview-260215": {
|
||||
"litellm_provider": "volcengine",
|
||||
"max_input_tokens": 256000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"source": "https://www.volcengine.com/docs/82379/1330310",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_tool_choice": false,
|
||||
"supports_vision": true,
|
||||
"tiered_pricing": [
|
||||
{
|
||||
"input_cost_per_token": 4.6e-07,
|
||||
"output_cost_per_token": 2.3e-06,
|
||||
"range": [
|
||||
0,
|
||||
32000.0
|
||||
]
|
||||
},
|
||||
{
|
||||
"input_cost_per_token": 7e-07,
|
||||
"output_cost_per_token": 3.5e-06,
|
||||
"range": [
|
||||
32000.0,
|
||||
128000.0
|
||||
]
|
||||
},
|
||||
{
|
||||
"input_cost_per_token": 1.4e-06,
|
||||
"output_cost_per_token": 7e-06,
|
||||
"range": [
|
||||
128000.0,
|
||||
256000.0
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
|
||||
Generated
+1
-1
@@ -8018,4 +8018,4 @@ utils = ["numpydoc"]
|
||||
[metadata]
|
||||
lock-version = "2.1"
|
||||
python-versions = ">=3.9,<4.0"
|
||||
content-hash = "eda34dfd8b35474beffee18893d6782c7b3d0d3d2c610f66237eb97176f43527"
|
||||
content-hash = "2cf958f1a04fd5f1ab0e5cfc33bdbf441b518ed6c82d0f2546bf64cd3d2f89be"
|
||||
|
||||
@@ -30,9 +30,11 @@ from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterat
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
from litellm.types.llms.openai import (
|
||||
ResponseCompletedEvent,
|
||||
ResponseFailedEvent,
|
||||
ResponseIncompleteEvent,
|
||||
ResponsesAPIResponse,
|
||||
ResponsesAPIStreamEvents,
|
||||
OutputTextDeltaEvent
|
||||
OutputTextDeltaEvent,
|
||||
)
|
||||
|
||||
|
||||
@@ -429,3 +431,155 @@ class TestBaseResponsesAPIStreamingIterator:
|
||||
mock_logging_obj.async_failure_handler.assert_not_called()
|
||||
mock_logging_obj.failure_handler.assert_not_called()
|
||||
|
||||
def test_process_chunk_response_failed_calls_failure_handler(self):
|
||||
"""
|
||||
Test that a RESPONSE_FAILED event routes to failure handlers,
|
||||
not success handlers. Failed responses represent genuine LLM-level
|
||||
errors and should be logged as failures.
|
||||
"""
|
||||
from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator
|
||||
|
||||
mock_response = Mock()
|
||||
mock_response.headers = {}
|
||||
mock_response.aiter_lines = Mock()
|
||||
mock_logging_obj = Mock(spec=LiteLLMLoggingObj)
|
||||
mock_logging_obj.model_call_details = {"litellm_params": {}}
|
||||
mock_logging_obj.async_failure_handler = Mock()
|
||||
mock_logging_obj.failure_handler = Mock()
|
||||
mock_logging_obj.async_success_handler = Mock()
|
||||
mock_logging_obj.success_handler = Mock()
|
||||
mock_config = Mock(spec=BaseResponsesAPIConfig)
|
||||
|
||||
mock_responses_api_response = Mock(spec=ResponsesAPIResponse)
|
||||
mock_responses_api_response.id = "resp_failed_123"
|
||||
mock_responses_api_response.error = {
|
||||
"type": "server_error",
|
||||
"message": "The model encountered an error",
|
||||
}
|
||||
mock_responses_api_response.usage = None
|
||||
|
||||
mock_failed_event = Mock(spec=ResponseFailedEvent)
|
||||
mock_failed_event.type = ResponsesAPIStreamEvents.RESPONSE_FAILED
|
||||
mock_failed_event.response = mock_responses_api_response
|
||||
|
||||
mock_config.transform_streaming_response.return_value = mock_failed_event
|
||||
|
||||
iterator = ResponsesAPIStreamingIterator(
|
||||
response=mock_response,
|
||||
model="gpt-4",
|
||||
responses_api_provider_config=mock_config,
|
||||
logging_obj=mock_logging_obj,
|
||||
litellm_metadata={"model_info": {"id": "model_123"}},
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
test_chunk_data = {
|
||||
"type": "response.failed",
|
||||
"response": {
|
||||
"id": "resp_failed_123",
|
||||
"error": {
|
||||
"type": "server_error",
|
||||
"message": "The model encountered an error",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
with patch.object(
|
||||
ResponsesAPIRequestUtils,
|
||||
"_update_responses_api_response_id_with_model_id",
|
||||
return_value=mock_responses_api_response,
|
||||
), patch(
|
||||
"litellm.responses.streaming_iterator.run_async_function"
|
||||
) as mock_run_async, patch(
|
||||
"litellm.responses.streaming_iterator.executor"
|
||||
) as mock_executor:
|
||||
result = iterator._process_chunk(json.dumps(test_chunk_data))
|
||||
|
||||
assert result is not None
|
||||
assert result.type == ResponsesAPIStreamEvents.RESPONSE_FAILED
|
||||
assert iterator.completed_response == result
|
||||
|
||||
# Failure handler should have been called via _handle_failure
|
||||
mock_run_async.assert_called_once()
|
||||
call_kwargs = mock_run_async.call_args
|
||||
assert (
|
||||
call_kwargs[1]["async_function"]
|
||||
== mock_logging_obj.async_failure_handler
|
||||
)
|
||||
|
||||
mock_executor.submit.assert_called_once()
|
||||
submit_args = mock_executor.submit.call_args
|
||||
assert submit_args[0][0] == mock_logging_obj.failure_handler
|
||||
|
||||
def test_process_chunk_response_incomplete_calls_success_handler(self):
|
||||
"""
|
||||
Test that a RESPONSE_INCOMPLETE event routes to success handlers.
|
||||
Incomplete responses (e.g. max_output_tokens reached) are still valid
|
||||
responses with usage data — analogous to finish_reason='length' in chat.
|
||||
"""
|
||||
from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator
|
||||
|
||||
mock_response = Mock()
|
||||
mock_response.headers = {}
|
||||
mock_response.aiter_lines = Mock()
|
||||
mock_logging_obj = Mock(spec=LiteLLMLoggingObj)
|
||||
mock_logging_obj.model_call_details = {"litellm_params": {}}
|
||||
mock_logging_obj.async_failure_handler = Mock()
|
||||
mock_logging_obj.failure_handler = Mock()
|
||||
mock_logging_obj.async_success_handler = Mock()
|
||||
mock_logging_obj.success_handler = Mock()
|
||||
mock_config = Mock(spec=BaseResponsesAPIConfig)
|
||||
|
||||
mock_responses_api_response = Mock(spec=ResponsesAPIResponse)
|
||||
mock_responses_api_response.id = "resp_incomplete_123"
|
||||
mock_responses_api_response.incomplete_details = {
|
||||
"reason": "max_output_tokens"
|
||||
}
|
||||
mock_responses_api_response.usage = None
|
||||
|
||||
mock_incomplete_event = Mock(spec=ResponseIncompleteEvent)
|
||||
mock_incomplete_event.type = ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE
|
||||
mock_incomplete_event.response = mock_responses_api_response
|
||||
|
||||
mock_config.transform_streaming_response.return_value = mock_incomplete_event
|
||||
|
||||
iterator = ResponsesAPIStreamingIterator(
|
||||
response=mock_response,
|
||||
model="gpt-4",
|
||||
responses_api_provider_config=mock_config,
|
||||
logging_obj=mock_logging_obj,
|
||||
litellm_metadata={"model_info": {"id": "model_123"}},
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
test_chunk_data = {
|
||||
"type": "response.incomplete",
|
||||
"response": {
|
||||
"id": "resp_incomplete_123",
|
||||
"incomplete_details": {"reason": "max_output_tokens"},
|
||||
},
|
||||
}
|
||||
|
||||
with patch.object(
|
||||
ResponsesAPIRequestUtils,
|
||||
"_update_responses_api_response_id_with_model_id",
|
||||
return_value=mock_responses_api_response,
|
||||
), patch(
|
||||
"asyncio.create_task"
|
||||
) as mock_create_task, patch(
|
||||
"litellm.responses.streaming_iterator.executor"
|
||||
) as mock_executor:
|
||||
result = iterator._process_chunk(json.dumps(test_chunk_data))
|
||||
|
||||
assert result is not None
|
||||
assert result.type == ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE
|
||||
assert iterator.completed_response == result
|
||||
|
||||
# Success handler should have been called (via _handle_logging_completed_response)
|
||||
mock_create_task.assert_called_once()
|
||||
mock_executor.submit.assert_called_once()
|
||||
|
||||
# Failure handlers should NOT have been called
|
||||
mock_logging_obj.async_failure_handler.assert_not_called()
|
||||
mock_logging_obj.failure_handler.assert_not_called()
|
||||
|
||||
|
||||
@@ -593,7 +593,7 @@ def test_datadog_static_methods():
|
||||
# Test tags format with default values
|
||||
assert (
|
||||
"env:unknown,service:litellm-server,version:unknown,HOSTNAME:"
|
||||
in get_datadog_tags()
|
||||
in ",".join(get_datadog_tags())
|
||||
)
|
||||
|
||||
# Test with custom environment variables
|
||||
@@ -631,7 +631,7 @@ def test_datadog_static_methods():
|
||||
# Test tags format with custom values
|
||||
expected_custom_tags = "env:production,service:custom-service,version:1.0.0,HOSTNAME:test-host,POD_NAME:pod-123"
|
||||
print("DataDogLogger._get_datadog_tags()", get_datadog_tags())
|
||||
assert get_datadog_tags() == expected_custom_tags
|
||||
assert ",".join(get_datadog_tags()) == expected_custom_tags
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -672,11 +672,11 @@ def test_get_datadog_tags():
|
||||
"""Test the _get_datadog_tags static method with various inputs"""
|
||||
# Test with no standard_logging_object and default env vars
|
||||
base_tags = get_datadog_tags()
|
||||
assert "env:" in base_tags
|
||||
assert "service:" in base_tags
|
||||
assert "version:" in base_tags
|
||||
assert "POD_NAME:" in base_tags
|
||||
assert "HOSTNAME:" in base_tags
|
||||
assert any("env:" in t for t in base_tags)
|
||||
assert any("service:" in t for t in base_tags)
|
||||
assert any("version:" in t for t in base_tags)
|
||||
assert any("POD_NAME:" in t for t in base_tags)
|
||||
assert any("HOSTNAME:" in t for t in base_tags)
|
||||
|
||||
# Test with custom env vars
|
||||
test_env = {
|
||||
@@ -705,12 +705,12 @@ def test_get_datadog_tags():
|
||||
# Test with empty request_tags
|
||||
standard_logging_obj["request_tags"] = []
|
||||
tags_empty_request = get_datadog_tags(standard_logging_obj)
|
||||
assert "request_tag:" not in tags_empty_request
|
||||
assert not any(t.startswith("request_tag:") for t in tags_empty_request)
|
||||
|
||||
# Test with None request_tags
|
||||
standard_logging_obj["request_tags"] = None
|
||||
tags_none_request = get_datadog_tags(standard_logging_obj)
|
||||
assert "request_tag:" not in tags_none_request
|
||||
assert not any(t.startswith("request_tag:") for t in tags_none_request)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@@ -2278,6 +2278,75 @@ async def test_post_call_failure_hook_auth_error_llm_api_route():
|
||||
mock_handle_logging.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"request_data, route, expected_call_type",
|
||||
[
|
||||
(
|
||||
{"model": "bad-model", "messages": [{"role": "user", "content": "hello"}]},
|
||||
"/v1/chat/completions",
|
||||
"acompletion",
|
||||
),
|
||||
(
|
||||
{"model": "bad-model", "prompt": "hello"},
|
||||
"/v1/completions",
|
||||
"atext_completion",
|
||||
),
|
||||
(
|
||||
{"model": "bad-model", "input": ["hello"]},
|
||||
"/v1/embeddings",
|
||||
"aembedding",
|
||||
),
|
||||
],
|
||||
)
|
||||
async def test_handle_logging_proxy_only_error_syncs_normalized_call_type(
|
||||
request_data, route, expected_call_type
|
||||
):
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
cache = DualCache()
|
||||
proxy_logging = ProxyLogging(user_api_key_cache=cache)
|
||||
captured_logging_obj = {}
|
||||
original_function_setup = litellm.utils.function_setup
|
||||
|
||||
def _capture_function_setup(*args, **kwargs):
|
||||
logging_obj, data = original_function_setup(*args, **kwargs)
|
||||
captured_logging_obj["logging_obj"] = logging_obj
|
||||
return logging_obj, data
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.utils.litellm.utils.function_setup",
|
||||
side_effect=_capture_function_setup,
|
||||
), patch.object(
|
||||
Logging, "async_failure_handler", new=AsyncMock(return_value=None)
|
||||
), patch.object(
|
||||
Logging, "failure_handler", return_value=None
|
||||
), patch(
|
||||
"litellm.proxy.utils.threading.Thread"
|
||||
) as mock_thread:
|
||||
mock_thread.return_value.start = Mock()
|
||||
|
||||
await proxy_logging._handle_logging_proxy_only_error(
|
||||
request_data=request_data,
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
api_key="test_key",
|
||||
user_id="test_user",
|
||||
token="test_token",
|
||||
request_route=route,
|
||||
),
|
||||
route=route,
|
||||
original_exception=HTTPException(status_code=400, detail="bad request"),
|
||||
)
|
||||
|
||||
logging_obj = captured_logging_obj["logging_obj"]
|
||||
assert logging_obj.call_type == expected_call_type
|
||||
assert logging_obj.model_call_details["call_type"] == expected_call_type
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_during_call_hook_parallel_execution():
|
||||
"""
|
||||
|
||||
@@ -1568,6 +1568,91 @@ def test_handle_clientside_credential_with_deployment_model_name(model_list):
|
||||
print("✓ _handle_clientside_credential test passed!")
|
||||
|
||||
|
||||
def test_sync_generic_api_call_preserves_requested_model_group_in_logs():
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "claude-sonnet-4-6",
|
||||
"litellm_params": {
|
||||
"model": "bedrock/global.anthropic.claude-sonnet-4-6",
|
||||
"aws_access_key_id": "test-access-key",
|
||||
"aws_secret_access_key": "test-secret-key",
|
||||
"aws_region_name": "us-west-2",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
try:
|
||||
captured_kwargs = {}
|
||||
|
||||
def mock_original_function(**kwargs):
|
||||
captured_kwargs.update(kwargs)
|
||||
return {"status": "ok"}
|
||||
|
||||
response = router._generic_api_call_with_fallbacks(
|
||||
model="claude-sonnet-4-6",
|
||||
original_function=mock_original_function,
|
||||
)
|
||||
|
||||
assert response == {"status": "ok"}
|
||||
assert (
|
||||
captured_kwargs["model"] == "bedrock/global.anthropic.claude-sonnet-4-6"
|
||||
)
|
||||
assert (
|
||||
captured_kwargs["litellm_metadata"]["model_group"] == "claude-sonnet-4-6"
|
||||
)
|
||||
assert (
|
||||
captured_kwargs["litellm_metadata"]["deployment"]
|
||||
== "bedrock/global.anthropic.claude-sonnet-4-6"
|
||||
)
|
||||
finally:
|
||||
router.discard()
|
||||
|
||||
|
||||
def test_sync_generic_api_call_uses_request_kwargs_for_deployment_selection():
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "regional-model",
|
||||
"litellm_params": {
|
||||
"model": "anthropic/us-model",
|
||||
"api_key": "test-api-key",
|
||||
"region_name": "us",
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "regional-model",
|
||||
"litellm_params": {
|
||||
"model": "anthropic/eu-model",
|
||||
"api_key": "test-api-key",
|
||||
"region_name": "eu",
|
||||
},
|
||||
},
|
||||
],
|
||||
enable_pre_call_checks=True,
|
||||
)
|
||||
|
||||
try:
|
||||
captured_kwargs = {}
|
||||
|
||||
def mock_original_function(**kwargs):
|
||||
captured_kwargs.update(kwargs)
|
||||
return {"status": "ok"}
|
||||
|
||||
response = router._generic_api_call_with_fallbacks(
|
||||
model="regional-model",
|
||||
original_function=mock_original_function,
|
||||
messages=[{"role": "user", "content": "Hello from Europe"}],
|
||||
allowed_model_region="eu",
|
||||
)
|
||||
|
||||
assert response == {"status": "ok"}
|
||||
assert captured_kwargs["model"] == "anthropic/eu-model"
|
||||
finally:
|
||||
router.discard()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"function_name, expected_metadata_key",
|
||||
[
|
||||
|
||||
@@ -44,7 +44,7 @@ class TestDatadogTagsRegression:
|
||||
assert "env:test-env" in tags_legacy
|
||||
assert "service:test-service" in tags_legacy
|
||||
# Verify NO team tag (should not invent one)
|
||||
assert "team:" not in tags_legacy
|
||||
assert not any(t.startswith("team:") for t in tags_legacy)
|
||||
|
||||
# Case 2: New feature (team info provided)
|
||||
payload_with_team = StandardLoggingPayload(
|
||||
|
||||
@@ -132,3 +132,57 @@ class TestLangsmithLoggerInit:
|
||||
assert (
|
||||
logger.sampling_rate >= 0.0
|
||||
), f"sampling_rate should be non-negative, got {logger.sampling_rate}"
|
||||
|
||||
|
||||
class TestLangsmithPrepareLogData:
|
||||
"""Regression test for #24001: _prepare_log_data must inject
|
||||
usage_metadata into outputs so LangSmith's Cost column is populated."""
|
||||
|
||||
@patch("asyncio.create_task")
|
||||
@patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": "1"}, clear=False)
|
||||
def test_outputs_contain_usage_metadata(self, mock_create_task):
|
||||
logger = LangsmithLogger(
|
||||
langsmith_api_key="test-key",
|
||||
langsmith_project="test-project",
|
||||
)
|
||||
|
||||
payload = {
|
||||
"id": "test-id",
|
||||
"response": {"choices": [{"message": {"content": "hi"}}]},
|
||||
"metadata": {},
|
||||
"startTime": 1.0,
|
||||
"endTime": 2.0,
|
||||
"request_tags": [],
|
||||
"error_str": None,
|
||||
"status": "success",
|
||||
"response_cost": 0.0042,
|
||||
"prompt_tokens": 100,
|
||||
"completion_tokens": 50,
|
||||
"total_tokens": 150,
|
||||
}
|
||||
|
||||
kwargs = {
|
||||
"litellm_params": {"metadata": {}},
|
||||
"standard_logging_object": payload,
|
||||
}
|
||||
|
||||
credentials = {
|
||||
"LANGSMITH_API_KEY": "test-key",
|
||||
"LANGSMITH_PROJECT": "test-project",
|
||||
"LANGSMITH_BASE_URL": "https://api.smith.langchain.com",
|
||||
}
|
||||
|
||||
data = logger._prepare_log_data(
|
||||
kwargs=kwargs,
|
||||
response_obj=None,
|
||||
start_time=1.0,
|
||||
end_time=2.0,
|
||||
credentials=credentials,
|
||||
)
|
||||
|
||||
assert "usage_metadata" in data["outputs"]
|
||||
um = data["outputs"]["usage_metadata"]
|
||||
assert um["total_cost"] == 0.0042
|
||||
assert um["input_tokens"] == 100
|
||||
assert um["output_tokens"] == 50
|
||||
assert um["total_tokens"] == 150
|
||||
|
||||
@@ -11,7 +11,8 @@ sys.path.insert(
|
||||
import time
|
||||
|
||||
from litellm.constants import SENTRY_DENYLIST, SENTRY_PII_DENYLIST
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
Logging as LitellmLogging
|
||||
from litellm.litellm_core_utils.litellm_logging import set_callbacks
|
||||
from litellm.types.utils import ModelResponse, TextCompletionResponse
|
||||
|
||||
@@ -139,7 +140,8 @@ def test_sentry_environment():
|
||||
|
||||
|
||||
def test_use_custom_pricing_for_model():
|
||||
from litellm.litellm_core_utils.litellm_logging import use_custom_pricing_for_model
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
use_custom_pricing_for_model
|
||||
|
||||
litellm_params = {
|
||||
"custom_llm_provider": "azure",
|
||||
@@ -154,7 +156,8 @@ def test_use_custom_pricing_for_model_via_litellm_metadata():
|
||||
Generic API call routes (/messages, /responses) store model_info
|
||||
under litellm_metadata, not metadata. Regression test for #23185.
|
||||
"""
|
||||
from litellm.litellm_core_utils.litellm_logging import use_custom_pricing_for_model
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
use_custom_pricing_for_model
|
||||
|
||||
litellm_params = {
|
||||
"litellm_metadata": {
|
||||
@@ -170,7 +173,8 @@ def test_use_custom_pricing_for_model_via_litellm_metadata():
|
||||
|
||||
def test_use_custom_pricing_not_detected_litellm_metadata_no_pricing():
|
||||
"""Should return False when litellm_metadata.model_info has no pricing keys."""
|
||||
from litellm.litellm_core_utils.litellm_logging import use_custom_pricing_for_model
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
use_custom_pricing_for_model
|
||||
|
||||
litellm_params = {
|
||||
"litellm_metadata": {
|
||||
@@ -186,7 +190,8 @@ def test_response_cost_calculator_uses_router_model_id_from_litellm_metadata():
|
||||
does not carry _hidden_params (e.g. ResponsesAPIResponse from /v1/responses
|
||||
streaming). Regression test for custom pricing on streaming responses."""
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
Logging as LiteLLMLoggingObj
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
|
||||
custom_model_id = "gpt-5-custom-pricing"
|
||||
@@ -256,6 +261,121 @@ def test_response_cost_calculator_uses_router_model_id_from_litellm_metadata():
|
||||
litellm.model_cost.pop(custom_model_id, None)
|
||||
|
||||
|
||||
class TestGetRouterModelId:
|
||||
"""Tests for the get_router_model_id helper method."""
|
||||
|
||||
def test_returns_id_from_litellm_metadata(self, logging_obj):
|
||||
"""Should extract model_info.id from litellm_metadata."""
|
||||
logging_obj.litellm_params = {
|
||||
"litellm_metadata": {
|
||||
"model_info": {"id": "custom-deploy-1"},
|
||||
},
|
||||
}
|
||||
assert logging_obj.get_router_model_id() == "custom-deploy-1"
|
||||
|
||||
def test_returns_id_from_metadata(self, logging_obj):
|
||||
"""Should fall back to metadata when litellm_metadata has no model_info."""
|
||||
logging_obj.litellm_params = {
|
||||
"metadata": {
|
||||
"model_info": {"id": "custom-deploy-2"},
|
||||
},
|
||||
}
|
||||
assert logging_obj.get_router_model_id() == "custom-deploy-2"
|
||||
|
||||
def test_prefers_litellm_metadata_over_metadata(self, logging_obj):
|
||||
"""litellm_metadata should take priority over metadata."""
|
||||
logging_obj.litellm_params = {
|
||||
"litellm_metadata": {
|
||||
"model_info": {"id": "from-litellm-meta"},
|
||||
},
|
||||
"metadata": {
|
||||
"model_info": {"id": "from-meta"},
|
||||
},
|
||||
}
|
||||
assert logging_obj.get_router_model_id() == "from-litellm-meta"
|
||||
|
||||
def test_returns_none_when_no_model_info(self, logging_obj):
|
||||
"""Should return None when no model_info is present."""
|
||||
logging_obj.litellm_params = {"api_base": ""}
|
||||
assert logging_obj.get_router_model_id() is None
|
||||
|
||||
def test_returns_none_when_no_litellm_params(self):
|
||||
"""Should return None when litellm_params is not set."""
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
Logging as LiteLLMLoggingObj
|
||||
|
||||
obj = LiteLLMLoggingObj(
|
||||
model="test",
|
||||
messages=[],
|
||||
stream=False,
|
||||
call_type="completion",
|
||||
start_time=time.time(),
|
||||
litellm_call_id="x",
|
||||
function_id="x",
|
||||
)
|
||||
# litellm_params exists but is empty by default
|
||||
assert obj.get_router_model_id() is None
|
||||
|
||||
|
||||
class TestAnthropicPassthroughCustomPricing:
|
||||
"""Verify the Anthropic pass-through handler forwards custom pricing."""
|
||||
|
||||
def test_completion_cost_receives_custom_pricing_args(self):
|
||||
"""_create_anthropic_response_logging_payload should pass
|
||||
custom_pricing and router_model_id to litellm.completion_cost
|
||||
when the logging object carries custom pricing in model_info."""
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import \
|
||||
AnthropicPassthroughLoggingHandler
|
||||
|
||||
logging_obj = LiteLLMLoggingObj(
|
||||
model="claude-sonnet-4-20250514",
|
||||
messages=[{"role": "user", "content": "Hi"}],
|
||||
stream=False,
|
||||
call_type="anthropic_messages",
|
||||
start_time=time.time(),
|
||||
litellm_call_id="test-456",
|
||||
function_id="test-fn",
|
||||
)
|
||||
logging_obj.update_environment_variables(
|
||||
model="claude-sonnet-4-20250514",
|
||||
user="",
|
||||
optional_params={},
|
||||
litellm_params={
|
||||
"api_base": "",
|
||||
"litellm_metadata": {
|
||||
"model_info": {
|
||||
"id": "claude-custom-pricing",
|
||||
"input_cost_per_token": 0.5,
|
||||
"output_cost_per_token": 1.5,
|
||||
},
|
||||
},
|
||||
},
|
||||
)
|
||||
logging_obj.model_call_details["custom_llm_provider"] = "anthropic"
|
||||
|
||||
mock_response = ModelResponse()
|
||||
mock_response.usage = {"prompt_tokens": 10, "completion_tokens": 5} # type: ignore
|
||||
|
||||
with patch("litellm.completion_cost", return_value=42.0) as mock_cost:
|
||||
AnthropicPassthroughLoggingHandler._create_anthropic_response_logging_payload(
|
||||
litellm_model_response=mock_response,
|
||||
model="claude-sonnet-4-20250514",
|
||||
kwargs={},
|
||||
start_time=time.time(),
|
||||
end_time=time.time(),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
mock_cost.assert_called_once()
|
||||
call_kwargs = mock_cost.call_args
|
||||
assert call_kwargs.kwargs.get("custom_pricing") is True
|
||||
assert call_kwargs.kwargs.get("router_model_id") == "claude-custom-pricing"
|
||||
|
||||
|
||||
class TestUpdateFromKwargs:
|
||||
"""Tests for the update_from_kwargs convenience wrapper."""
|
||||
|
||||
@@ -321,9 +441,8 @@ class TestUpdateFromKwargs:
|
||||
|
||||
def test_custom_pricing_detected_via_litellm_metadata(self, logging_obj):
|
||||
"""Custom pricing in litellm_metadata.model_info should set custom_pricing flag."""
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
use_custom_pricing_for_model,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
use_custom_pricing_for_model
|
||||
|
||||
lm_meta = {
|
||||
"model_info": {
|
||||
@@ -382,7 +501,8 @@ async def test_datadog_logger_not_shadowed_by_llm_obs(monkeypatch):
|
||||
monkeypatch.setenv("DD_SITE", "us5.datadoghq.com")
|
||||
|
||||
from litellm.integrations.datadog.datadog import DataDogLogger
|
||||
from litellm.integrations.datadog.datadog_llm_obs import DataDogLLMObsLogger
|
||||
from litellm.integrations.datadog.datadog_llm_obs import \
|
||||
DataDogLLMObsLogger
|
||||
from litellm.litellm_core_utils import litellm_logging as logging_module
|
||||
|
||||
logging_module._in_memory_loggers.clear()
|
||||
@@ -423,7 +543,8 @@ async def test_logfire_logger_accepts_env_vars_for_base_url(monkeypatch):
|
||||
) # no trailing slash on purpose
|
||||
|
||||
# Import after env vars are set (important if module-level caching exists)
|
||||
from litellm.integrations.opentelemetry import OpenTelemetry # logger class
|
||||
from litellm.integrations.opentelemetry import \
|
||||
OpenTelemetry # logger class
|
||||
from litellm.litellm_core_utils import litellm_logging as logging_module
|
||||
|
||||
logging_module._in_memory_loggers.clear()
|
||||
@@ -752,7 +873,8 @@ def test_success_handler_runs_guardrail_logging_hook_when_enabled(logging_obj):
|
||||
|
||||
|
||||
def test_get_user_agent_tags():
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
StandardLoggingPayloadSetup
|
||||
|
||||
tags = StandardLoggingPayloadSetup._get_user_agent_tags(
|
||||
proxy_server_request={
|
||||
@@ -767,7 +889,8 @@ def test_get_user_agent_tags():
|
||||
|
||||
|
||||
def test_get_request_tags():
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
StandardLoggingPayloadSetup
|
||||
|
||||
tags = StandardLoggingPayloadSetup._get_request_tags(
|
||||
litellm_params={"metadata": {"tags": ["test-tag"]}},
|
||||
@@ -794,7 +917,8 @@ def test_get_request_tags_from_metadata_and_litellm_metadata():
|
||||
4. No tags in either
|
||||
5. None values for metadata/litellm_metadata
|
||||
"""
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
StandardLoggingPayloadSetup
|
||||
|
||||
# Test case 1: Tags in metadata only
|
||||
tags = StandardLoggingPayloadSetup._get_request_tags(
|
||||
@@ -875,7 +999,8 @@ def test_get_request_tags_does_not_mutate_original_tags():
|
||||
would cause User-Agent tags to be duplicated because the function was mutating
|
||||
the original tags list instead of creating a copy.
|
||||
"""
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
StandardLoggingPayloadSetup
|
||||
|
||||
# Create metadata with original tags
|
||||
original_tags = ["custom-tag-1", "custom-tag-2"]
|
||||
@@ -935,7 +1060,8 @@ def test_get_request_tags_does_not_mutate_original_tags():
|
||||
def test_get_extra_header_tags():
|
||||
"""Test the _get_extra_header_tags method with various scenarios."""
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
StandardLoggingPayloadSetup
|
||||
|
||||
# Store original value to restore later
|
||||
original_extra_headers = getattr(litellm, "extra_spend_tag_headers", None)
|
||||
@@ -1156,7 +1282,8 @@ async def test_e2e_generate_cold_storage_object_key_successful():
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
StandardLoggingPayloadSetup
|
||||
|
||||
# Create test data
|
||||
start_time = datetime(2025, 1, 15, 10, 30, 45, 123456, timezone.utc)
|
||||
@@ -1198,7 +1325,8 @@ async def test_e2e_generate_cold_storage_object_key_with_custom_logger_s3_path()
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
StandardLoggingPayloadSetup
|
||||
|
||||
# Create test data
|
||||
start_time = datetime(2025, 1, 15, 10, 30, 45, 123456, timezone.utc)
|
||||
@@ -1249,7 +1377,8 @@ async def test_e2e_generate_cold_storage_object_key_with_logger_no_s3_path():
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
StandardLoggingPayloadSetup
|
||||
|
||||
# Create test data
|
||||
start_time = datetime(2025, 1, 15, 10, 30, 45, 123456, timezone.utc)
|
||||
@@ -1296,7 +1425,8 @@ async def test_e2e_generate_cold_storage_object_key_not_configured():
|
||||
from unittest.mock import patch
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
StandardLoggingPayloadSetup
|
||||
|
||||
# Create test data
|
||||
start_time = datetime(2025, 1, 15, 10, 30, 45, 123456, timezone.utc)
|
||||
@@ -1320,7 +1450,8 @@ def test_get_final_response_obj_with_empty_response_obj_and_list_init():
|
||||
|
||||
When response_obj is empty (falsy), the method should return init_response_obj if it's a list.
|
||||
"""
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
StandardLoggingPayloadSetup
|
||||
|
||||
# Create test objects
|
||||
class TestObject1:
|
||||
@@ -1356,7 +1487,8 @@ def test_get_usage_as_dict():
|
||||
"""
|
||||
Test get_usage_as_dict returns usage as plain dict from response_obj or combined_usage_object.
|
||||
"""
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
StandardLoggingPayloadSetup
|
||||
from litellm.types.utils import Usage
|
||||
|
||||
# Test case 1: None response_obj returns empty usage dict
|
||||
@@ -1394,7 +1526,8 @@ def test_append_system_prompt_messages():
|
||||
"""
|
||||
Test append_system_prompt_messages prepends system message from kwargs to messages list.
|
||||
"""
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
StandardLoggingPayloadSetup
|
||||
|
||||
# Test case 1: system in kwargs with existing messages
|
||||
kwargs = {"system": "You are a helpful assistant"}
|
||||
@@ -1465,7 +1598,8 @@ async def test_async_success_handler_sets_standard_logging_object_for_pass_throu
|
||||
from datetime import datetime
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
Logging as LiteLLMLoggingObj
|
||||
from litellm.types.utils import StandardPassThroughResponseObject
|
||||
|
||||
# Create a logging object for a pass-through endpoint
|
||||
@@ -1546,7 +1680,8 @@ async def test_async_success_handler_prevents_reprocessing_for_pass_through_endp
|
||||
from datetime import datetime
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
Logging as LiteLLMLoggingObj
|
||||
from litellm.types.utils import StandardPassThroughResponseObject
|
||||
|
||||
# Create a logging object for a pass-through endpoint
|
||||
@@ -1622,7 +1757,8 @@ async def test_async_success_handler_sets_standard_logging_object_for_streaming_
|
||||
from datetime import datetime
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
Logging as LiteLLMLoggingObj
|
||||
from litellm.types.utils import StandardPassThroughResponseObject
|
||||
|
||||
# Create a logging object for a streaming pass-through endpoint
|
||||
@@ -1678,7 +1814,8 @@ def test_get_error_information_error_code_priority():
|
||||
Test get_error_information prioritizes 'code' attribute over 'status_code' attribute
|
||||
and handles edge cases like empty strings and "None" string values.
|
||||
"""
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
StandardLoggingPayloadSetup
|
||||
|
||||
# Test case 1: Exception with 'code' attribute (ProxyException style)
|
||||
class ProxyException(Exception):
|
||||
@@ -1871,7 +2008,8 @@ async def test_async_success_handler_preserves_response_cost_for_pass_through_en
|
||||
by pass-through handlers (Gemini/Vertex)."""
|
||||
from datetime import datetime
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
Logging as LiteLLMLoggingObj
|
||||
from litellm.types.utils import ModelResponse, Usage
|
||||
|
||||
logging_obj = LiteLLMLoggingObj(
|
||||
|
||||
@@ -1340,6 +1340,121 @@ def test_is_chunk_non_empty_with_valid_tool_calls(
|
||||
)
|
||||
|
||||
|
||||
def _make_chunk(content: Optional[str]) -> ModelResponseStream:
|
||||
return ModelResponseStream(
|
||||
id="test",
|
||||
created=1741037890,
|
||||
model="test-model",
|
||||
choices=[StreamingChoices(index=0, delta=Delta(content=content))],
|
||||
)
|
||||
|
||||
|
||||
def _build_chunks(pattern: list[str], N: int) -> list[ModelResponseStream]:
|
||||
"""
|
||||
Build a list of chunks based on a pattern specification.
|
||||
"""
|
||||
chunks = []
|
||||
for i, p in enumerate(pattern):
|
||||
if p == "same":
|
||||
chunks.append(_make_chunk("same_chunk"))
|
||||
elif p == "diff":
|
||||
chunks.append(_make_chunk(f"chunk_{i}"))
|
||||
else:
|
||||
chunks.append(_make_chunk(p))
|
||||
return chunks
|
||||
|
||||
_REPETITION_TEST_CASES = [
|
||||
# Basic cases
|
||||
pytest.param(
|
||||
["same"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT,
|
||||
True,
|
||||
id="all_identical_raises",
|
||||
),
|
||||
pytest.param(
|
||||
["same"] * (litellm.REPEATED_STREAMING_CHUNK_LIMIT - 1),
|
||||
False,
|
||||
id="below_threshold_no_raise",
|
||||
),
|
||||
pytest.param(
|
||||
[None] * litellm.REPEATED_STREAMING_CHUNK_LIMIT,
|
||||
False,
|
||||
id="none_content_no_raise",
|
||||
),
|
||||
pytest.param(
|
||||
[""] * litellm.REPEATED_STREAMING_CHUNK_LIMIT,
|
||||
False,
|
||||
id="empty_content_no_raise",
|
||||
),
|
||||
# Short content (len <= 2) should not raise
|
||||
pytest.param(
|
||||
["##"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT,
|
||||
False,
|
||||
id="short_content_2chars_no_raise",
|
||||
),
|
||||
pytest.param(
|
||||
["{"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT,
|
||||
False,
|
||||
id="short_content_1char_no_raise",
|
||||
),
|
||||
pytest.param(
|
||||
["ab"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT,
|
||||
False,
|
||||
id="short_content_2chars_ab_no_raise",
|
||||
),
|
||||
# All different chunks
|
||||
pytest.param(
|
||||
["diff"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT,
|
||||
False,
|
||||
id="all_different_no_raise",
|
||||
),
|
||||
# One chunk different at various positions
|
||||
pytest.param(
|
||||
["different_first"] + ["same"] * (litellm.REPEATED_STREAMING_CHUNK_LIMIT - 1),
|
||||
False,
|
||||
id="first_chunk_different_no_raise",
|
||||
),
|
||||
pytest.param(
|
||||
["same"] * (litellm.REPEATED_STREAMING_CHUNK_LIMIT - 1) + ["different_last"],
|
||||
False,
|
||||
id="last_chunk_different_no_raise",
|
||||
),
|
||||
pytest.param(
|
||||
["same"] * (litellm.REPEATED_STREAMING_CHUNK_LIMIT // 2 + 1) + ["different_mid"] + ["same"] * (litellm.REPEATED_STREAMING_CHUNK_LIMIT - litellm.REPEATED_STREAMING_CHUNK_LIMIT // 2 + 1),
|
||||
False,
|
||||
id="middle_chunk_different_no_raise",
|
||||
),
|
||||
pytest.param(
|
||||
["same"] * (litellm.REPEATED_STREAMING_CHUNK_LIMIT - 2) + ["diff", "diff"],
|
||||
False,
|
||||
id="last_two_different_no_raise",
|
||||
),
|
||||
pytest.param(
|
||||
["diff"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT + ["same"] * litellm.REPEATED_STREAMING_CHUNK_LIMIT + ["diff"],
|
||||
True,
|
||||
id="in_between_same_and_diff_raise",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("chunks_pattern,should_raise", _REPETITION_TEST_CASES)
|
||||
def test_raise_on_model_repetition(
|
||||
initialized_custom_stream_wrapper: CustomStreamWrapper,
|
||||
chunks_pattern: list,
|
||||
should_raise: bool,
|
||||
):
|
||||
wrapper = initialized_custom_stream_wrapper
|
||||
chunks = _build_chunks(chunks_pattern, len(chunks_pattern))
|
||||
|
||||
if should_raise:
|
||||
with pytest.raises(litellm.InternalServerError) as exc_info:
|
||||
for chunk in chunks:
|
||||
wrapper.chunks.append(chunk)
|
||||
wrapper.raise_on_model_repetition()
|
||||
assert "repeating the same chunk" in str(exc_info.value)
|
||||
else:
|
||||
for chunk in chunks:
|
||||
wrapper.chunks.append(chunk)
|
||||
wrapper.raise_on_model_repetition()
|
||||
def test_usage_chunk_after_finish_reason_updates_hidden_params(logging_obj):
|
||||
"""
|
||||
Test that provider-reported usage from a post-finish_reason chunk
|
||||
|
||||
@@ -6,6 +6,7 @@ from litellm.types.llms.openai import (
|
||||
ChatCompletionToolCallChunk,
|
||||
ChatCompletionToolCallFunctionChunk,
|
||||
)
|
||||
from litellm.types.responses.main import OutputCodeInterpreterCall
|
||||
|
||||
|
||||
def test_redacted_thinking_content_block_delta():
|
||||
@@ -479,14 +480,22 @@ def test_partial_json_chunk_accumulation():
|
||||
# First partial chunk should return None (still accumulating)
|
||||
result1 = iterator._parse_sse_data(f"data:{partial_chunk_1}")
|
||||
assert result1 is None, "First partial chunk should return None while accumulating"
|
||||
assert iterator.chunk_type == "accumulated_json", "Should switch to accumulated_json mode"
|
||||
assert iterator.accumulated_json == partial_chunk_1, "Should have accumulated first part"
|
||||
assert (
|
||||
iterator.chunk_type == "accumulated_json"
|
||||
), "Should switch to accumulated_json mode"
|
||||
assert (
|
||||
iterator.accumulated_json == partial_chunk_1
|
||||
), "Should have accumulated first part"
|
||||
|
||||
# Second partial chunk should complete the JSON and return a parsed result
|
||||
result2 = iterator._parse_sse_data(f"data:{partial_chunk_2}")
|
||||
assert result2 is not None, "Second chunk should return parsed result"
|
||||
assert iterator.accumulated_json == "", "Buffer should be cleared after successful parse"
|
||||
assert result2.choices[0].delta.content == "Hello", f"Expected 'Hello', got '{result2.choices[0].delta.content}'"
|
||||
assert (
|
||||
iterator.accumulated_json == ""
|
||||
), "Buffer should be cleared after successful parse"
|
||||
assert (
|
||||
result2.choices[0].delta.content == "Hello"
|
||||
), f"Expected 'Hello', got '{result2.choices[0].delta.content}'"
|
||||
|
||||
|
||||
def test_complete_json_chunk_no_accumulation():
|
||||
@@ -503,7 +512,9 @@ def test_complete_json_chunk_no_accumulation():
|
||||
assert result is not None, "Complete chunk should return parsed result immediately"
|
||||
assert iterator.chunk_type == "valid_json", "Should remain in valid_json mode"
|
||||
assert iterator.accumulated_json == "", "Buffer should remain empty"
|
||||
assert result.choices[0].delta.content == "Hello", f"Expected 'Hello', got '{result.choices[0].delta.content}'"
|
||||
assert (
|
||||
result.choices[0].delta.content == "Hello"
|
||||
), f"Expected 'Hello', got '{result.choices[0].delta.content}'"
|
||||
|
||||
|
||||
def test_multiple_partial_chunks_accumulation():
|
||||
@@ -620,7 +631,9 @@ def test_web_search_tool_result_no_extra_tool_calls():
|
||||
# Should have exactly 2 tool calls:
|
||||
# 1. From content_block_start (server_tool_use) with id and name
|
||||
# 2. From content_block_delta with the actual query
|
||||
assert len(tool_calls_emitted) == 2, f"Expected 2 tool calls, got {len(tool_calls_emitted)}"
|
||||
assert (
|
||||
len(tool_calls_emitted) == 2
|
||||
), f"Expected 2 tool calls, got {len(tool_calls_emitted)}"
|
||||
|
||||
# First tool call should have the id and name
|
||||
assert tool_calls_emitted[0]["id"] == "srvtoolu_01ABC123"
|
||||
@@ -722,7 +735,10 @@ def test_web_search_tool_result_captured_in_provider_specific_fields():
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "input_json_delta", "partial_json": '{"query": "otter facts"}'},
|
||||
"delta": {
|
||||
"type": "input_json_delta",
|
||||
"partial_json": '{"query": "otter facts"}',
|
||||
},
|
||||
},
|
||||
# 4. content_block_stop for server_tool_use
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
@@ -822,7 +838,10 @@ def test_web_fetch_tool_result_captured_in_provider_specific_fields():
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "input_json_delta", "partial_json": '{"url": "https://example.com"}'},
|
||||
"delta": {
|
||||
"type": "input_json_delta",
|
||||
"partial_json": '{"url": "https://example.com"}',
|
||||
},
|
||||
},
|
||||
# 4. content_block_stop for server_tool_use
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
@@ -946,7 +965,7 @@ def test_web_fetch_tool_result_no_extra_tool_calls():
|
||||
def test_container_in_provider_specific_fields_streaming():
|
||||
"""
|
||||
Test that container is captured in provider_specific_fields for streaming responses.
|
||||
|
||||
|
||||
When container with skills is used, the container field should be present in
|
||||
the provider_specific_fields of the message_delta chunk.
|
||||
"""
|
||||
@@ -1025,7 +1044,9 @@ def test_container_in_provider_specific_fields_streaming():
|
||||
]
|
||||
|
||||
# Verify container was captured
|
||||
assert container_field is not None, "container should be captured in provider_specific_fields"
|
||||
assert (
|
||||
container_field is not None
|
||||
), "container should be captured in provider_specific_fields"
|
||||
assert (
|
||||
container_field["id"] == "container_011CW9hA9zpZ8xD3bjjShy4p"
|
||||
), "container id should match"
|
||||
@@ -1033,18 +1054,14 @@ def test_container_in_provider_specific_fields_streaming():
|
||||
container_field["expires_at"] == "2025-12-16T04:57:16.913181Z"
|
||||
), "expires_at should match"
|
||||
assert len(container_field["skills"]) == 1, "Should have 1 skill"
|
||||
assert (
|
||||
container_field["skills"][0]["skill_id"] == "pptx"
|
||||
), "skill_id should be pptx"
|
||||
assert (
|
||||
container_field["skills"][0]["version"] == "20251013"
|
||||
), "version should match"
|
||||
assert container_field["skills"][0]["skill_id"] == "pptx", "skill_id should be pptx"
|
||||
assert container_field["skills"][0]["version"] == "20251013", "version should match"
|
||||
|
||||
|
||||
def test_container_in_provider_specific_fields_non_streaming():
|
||||
"""
|
||||
Test that container is captured in provider_specific_fields for non-streaming responses.
|
||||
|
||||
|
||||
When container with skills is used in non-streaming, the container field should be
|
||||
present in the provider_specific_fields of the response.
|
||||
"""
|
||||
@@ -1106,7 +1123,7 @@ def test_container_in_provider_specific_fields_non_streaming():
|
||||
def test_container_absent_when_not_provided():
|
||||
"""
|
||||
Test that container is not added to provider_specific_fields when not provided.
|
||||
|
||||
|
||||
This ensures we don't add empty or None container fields.
|
||||
"""
|
||||
iterator = ModelResponseIterator(
|
||||
@@ -1133,3 +1150,434 @@ def test_container_absent_when_not_provided():
|
||||
assert (
|
||||
"container" not in model_response.choices[0].delta.provider_specific_fields
|
||||
), "container should not be present when not provided in delta"
|
||||
|
||||
|
||||
def test_streaming_code_execution_produces_code_interpreter_results():
|
||||
"""
|
||||
Test that bash_code_execution_tool_result content blocks in streaming
|
||||
produce code_interpreter_results in provider_specific_fields, so the
|
||||
Responses API layer can use them without Anthropic-specific knowledge.
|
||||
"""
|
||||
|
||||
chunks = [
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": "msg_01XYZ",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [],
|
||||
"usage": {"input_tokens": 100, "output_tokens": 1},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 0,
|
||||
"content_block": {
|
||||
"type": "text",
|
||||
"text": "",
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {"type": "text_delta", "text": "Running code..."},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 1,
|
||||
"content_block": {
|
||||
"type": "server_tool_use",
|
||||
"id": "srvtoolu_01ABC",
|
||||
"name": "bash_code_execution",
|
||||
"input": {"command": "echo hello"},
|
||||
},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 1},
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 2,
|
||||
"content_block": {
|
||||
"type": "bash_code_execution_tool_result",
|
||||
"tool_use_id": "srvtoolu_01ABC",
|
||||
"content": {
|
||||
"type": "bash_code_execution_result",
|
||||
"stdout": "hello\n",
|
||||
"stderr": "",
|
||||
"return_code": 0,
|
||||
},
|
||||
},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 2},
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn"},
|
||||
"usage": {"output_tokens": 50},
|
||||
},
|
||||
]
|
||||
|
||||
iterator = ModelResponseIterator(None, sync_stream=True)
|
||||
|
||||
found_code_interpreter_results = False
|
||||
for chunk in chunks:
|
||||
parsed = iterator.chunk_parser(chunk)
|
||||
psf = None
|
||||
if parsed.choices and parsed.choices[0].delta:
|
||||
psf = getattr(parsed.choices[0].delta, "provider_specific_fields", None)
|
||||
if psf and "code_interpreter_results" in psf:
|
||||
found_code_interpreter_results = True
|
||||
results = psf["code_interpreter_results"]
|
||||
assert len(results) == 1
|
||||
assert isinstance(results[0], OutputCodeInterpreterCall)
|
||||
assert results[0].type == "code_interpreter_call"
|
||||
assert results[0].id == "srvtoolu_01ABC"
|
||||
assert results[0].code == "echo hello"
|
||||
assert results[0].outputs is not None
|
||||
assert len(results[0].outputs) == 1
|
||||
assert results[0].outputs[0].logs == "hello\n"
|
||||
|
||||
assert found_code_interpreter_results, (
|
||||
"code_interpreter_results should appear in provider_specific_fields "
|
||||
"when bash_code_execution_tool_result is streamed"
|
||||
)
|
||||
|
||||
|
||||
def test_streaming_multiple_code_executions_no_duplicates():
|
||||
"""
|
||||
Test that multiple code executions in a single streaming response emit
|
||||
cumulative code_interpreter_results on each chunk (matching stream_chunk_builder's
|
||||
"last value wins" contract). The final emission must contain ALL results.
|
||||
"""
|
||||
chunks = [
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": "msg_01XYZ",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [],
|
||||
"usage": {"input_tokens": 100, "output_tokens": 1},
|
||||
},
|
||||
},
|
||||
# First code execution
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 0,
|
||||
"content_block": {
|
||||
"type": "server_tool_use",
|
||||
"id": "srvtoolu_01AAA",
|
||||
"name": "bash_code_execution",
|
||||
"input": {"command": "echo first"},
|
||||
},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 1,
|
||||
"content_block": {
|
||||
"type": "bash_code_execution_tool_result",
|
||||
"tool_use_id": "srvtoolu_01AAA",
|
||||
"content": {
|
||||
"type": "bash_code_execution_result",
|
||||
"stdout": "first\n",
|
||||
"stderr": "",
|
||||
"return_code": 0,
|
||||
},
|
||||
},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 1},
|
||||
# Second code execution
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 2,
|
||||
"content_block": {
|
||||
"type": "server_tool_use",
|
||||
"id": "srvtoolu_01BBB",
|
||||
"name": "bash_code_execution",
|
||||
"input": {"command": "echo second"},
|
||||
},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 2},
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 3,
|
||||
"content_block": {
|
||||
"type": "bash_code_execution_tool_result",
|
||||
"tool_use_id": "srvtoolu_01BBB",
|
||||
"content": {
|
||||
"type": "bash_code_execution_result",
|
||||
"stdout": "second\n",
|
||||
"stderr": "",
|
||||
"return_code": 0,
|
||||
},
|
||||
},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 3},
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn"},
|
||||
"usage": {"output_tokens": 50},
|
||||
},
|
||||
]
|
||||
|
||||
iterator = ModelResponseIterator(None, sync_stream=True)
|
||||
|
||||
# Collect each emission of code_interpreter_results
|
||||
emissions = []
|
||||
for chunk in chunks:
|
||||
parsed = iterator.chunk_parser(chunk)
|
||||
psf = None
|
||||
if parsed.choices and parsed.choices[0].delta:
|
||||
psf = getattr(parsed.choices[0].delta, "provider_specific_fields", None)
|
||||
if psf and "code_interpreter_results" in psf:
|
||||
emissions.append(psf["code_interpreter_results"])
|
||||
|
||||
# Should have 2 emissions (one per tool_result block)
|
||||
assert len(emissions) == 2, f"Expected 2 emissions, got {len(emissions)}"
|
||||
|
||||
# First emission: cumulative list with 1 result
|
||||
assert len(emissions[0]) == 1
|
||||
assert emissions[0][0].id == "srvtoolu_01AAA"
|
||||
assert emissions[0][0].code == "echo first"
|
||||
assert emissions[0][0].outputs[0].logs == "first\n"
|
||||
|
||||
# Second (final) emission: cumulative list with BOTH results
|
||||
# This is what stream_chunk_builder will pick as "last value wins"
|
||||
assert len(emissions[1]) == 2, (
|
||||
f"Expected final emission to have 2 results, got {len(emissions[1])}. "
|
||||
f"IDs: {[r.id for r in emissions[1]]}"
|
||||
)
|
||||
assert emissions[1][0].id == "srvtoolu_01AAA"
|
||||
assert emissions[1][0].code == "echo first"
|
||||
assert emissions[1][0].outputs[0].logs == "first\n"
|
||||
assert emissions[1][1].id == "srvtoolu_01BBB"
|
||||
assert emissions[1][1].code == "echo second"
|
||||
assert emissions[1][1].outputs[0].logs == "second\n"
|
||||
|
||||
|
||||
def test_streaming_code_execution_input_assembled_from_deltas():
|
||||
"""
|
||||
In real Anthropic streaming, content_block_start for server_tool_use has
|
||||
input: {}. The actual input arrives via input_json_delta deltas and must
|
||||
be assembled at content_block_stop so the code field is populated.
|
||||
|
||||
This test uses realistic chunk shapes (empty input in start, partial JSON
|
||||
in deltas) to exercise the input assembly path.
|
||||
"""
|
||||
chunks = [
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": "msg_01XYZ",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [],
|
||||
"usage": {"input_tokens": 100, "output_tokens": 1},
|
||||
},
|
||||
},
|
||||
# server_tool_use with empty input (real streaming behaviour)
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 0,
|
||||
"content_block": {
|
||||
"type": "server_tool_use",
|
||||
"id": "srvtoolu_01AAA",
|
||||
"name": "code_execution",
|
||||
"input": {},
|
||||
},
|
||||
},
|
||||
# Input arrives via deltas, split across two chunks
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {
|
||||
"type": "input_json_delta",
|
||||
"partial_json": '{"comma',
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {
|
||||
"type": "input_json_delta",
|
||||
"partial_json": 'nd": "echo hello"}',
|
||||
},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
# Tool result
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 1,
|
||||
"content_block": {
|
||||
"type": "bash_code_execution_tool_result",
|
||||
"tool_use_id": "srvtoolu_01AAA",
|
||||
"content": {
|
||||
"type": "bash_code_execution_result",
|
||||
"stdout": "hello\n",
|
||||
"stderr": "",
|
||||
"return_code": 0,
|
||||
},
|
||||
},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 1},
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn"},
|
||||
"usage": {"output_tokens": 50},
|
||||
},
|
||||
]
|
||||
|
||||
iterator = ModelResponseIterator(None, sync_stream=True)
|
||||
|
||||
code_results = None
|
||||
for chunk in chunks:
|
||||
parsed = iterator.chunk_parser(chunk)
|
||||
psf = None
|
||||
if parsed.choices and parsed.choices[0].delta:
|
||||
psf = getattr(parsed.choices[0].delta, "provider_specific_fields", None)
|
||||
if psf and "code_interpreter_results" in psf:
|
||||
code_results = psf["code_interpreter_results"]
|
||||
|
||||
# The code field must contain the assembled input, not be empty
|
||||
assert code_results is not None, "No code_interpreter_results emitted"
|
||||
assert len(code_results) == 1
|
||||
assert code_results[0].id == "srvtoolu_01AAA"
|
||||
assert code_results[0].code == "echo hello"
|
||||
assert code_results[0].outputs[0].logs == "hello\n"
|
||||
|
||||
|
||||
def test_empty_output_produces_null_outputs():
|
||||
"""
|
||||
When both stdout and stderr are empty, outputs should be None
|
||||
(matching OpenAI's native behavior) rather than [{logs: ""}].
|
||||
"""
|
||||
chunks = [
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": "msg_01XYZ",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [],
|
||||
"usage": {"input_tokens": 100, "output_tokens": 1},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 0,
|
||||
"content_block": {
|
||||
"type": "server_tool_use",
|
||||
"id": "srvtoolu_01AAA",
|
||||
"name": "bash_code_execution",
|
||||
"input": {"command": "true"},
|
||||
},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 1,
|
||||
"content_block": {
|
||||
"type": "bash_code_execution_tool_result",
|
||||
"tool_use_id": "srvtoolu_01AAA",
|
||||
"content": {
|
||||
"type": "bash_code_execution_result",
|
||||
"stdout": "",
|
||||
"stderr": "",
|
||||
"return_code": 0,
|
||||
},
|
||||
},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 1},
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn"},
|
||||
"usage": {"output_tokens": 50},
|
||||
},
|
||||
]
|
||||
|
||||
iterator = ModelResponseIterator(None, sync_stream=True)
|
||||
|
||||
code_results = None
|
||||
for chunk in chunks:
|
||||
parsed = iterator.chunk_parser(chunk)
|
||||
psf = None
|
||||
if parsed.choices and parsed.choices[0].delta:
|
||||
psf = getattr(parsed.choices[0].delta, "provider_specific_fields", None)
|
||||
if psf and "code_interpreter_results" in psf:
|
||||
code_results = psf["code_interpreter_results"]
|
||||
|
||||
assert code_results is not None, "No code_interpreter_results emitted"
|
||||
assert len(code_results) == 1
|
||||
assert code_results[0].id == "srvtoolu_01AAA"
|
||||
assert (
|
||||
code_results[0].outputs is None
|
||||
), f"Expected outputs=None for empty execution, got {code_results[0].outputs}"
|
||||
|
||||
|
||||
def test_non_bash_tool_result_skipped():
|
||||
"""
|
||||
Tool result types other than bash_code_execution_tool_result (e.g.
|
||||
text_editor_code_execution_tool_result) should be skipped and NOT
|
||||
produce code_interpreter_call items.
|
||||
"""
|
||||
chunks = [
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": "msg_01XYZ",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [],
|
||||
"usage": {"input_tokens": 100, "output_tokens": 1},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 0,
|
||||
"content_block": {
|
||||
"type": "server_tool_use",
|
||||
"id": "srvtoolu_01AAA",
|
||||
"name": "text_editor",
|
||||
"input": {"command": "view", "path": "/tmp/test.py"},
|
||||
},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
# text_editor result — should NOT become a code_interpreter_call
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 1,
|
||||
"content_block": {
|
||||
"type": "text_editor_code_execution_tool_result",
|
||||
"tool_use_id": "srvtoolu_01AAA",
|
||||
"content": [
|
||||
{"type": "text", "text": "file contents here"},
|
||||
],
|
||||
},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 1},
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn"},
|
||||
"usage": {"output_tokens": 50},
|
||||
},
|
||||
]
|
||||
|
||||
iterator = ModelResponseIterator(None, sync_stream=True)
|
||||
|
||||
code_results = None
|
||||
for chunk in chunks:
|
||||
parsed = iterator.chunk_parser(chunk)
|
||||
psf = None
|
||||
if parsed.choices and parsed.choices[0].delta:
|
||||
psf = getattr(parsed.choices[0].delta, "provider_specific_fields", None)
|
||||
if psf and "code_interpreter_results" in psf:
|
||||
code_results = psf["code_interpreter_results"]
|
||||
|
||||
# code_interpreter_results should be emitted but empty (no bash results)
|
||||
assert (
|
||||
code_results is not None
|
||||
), "Expected code_interpreter_results key to be emitted"
|
||||
assert (
|
||||
len(code_results) == 0
|
||||
), f"Expected 0 code_interpreter_results for text_editor result, got {len(code_results)}"
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,268 @@
|
||||
"""
|
||||
Tests for the Responses API _extract_tool_result_output_items path,
|
||||
the non-streaming _hidden_params propagation of code_interpreter_results,
|
||||
and mock end-to-end streaming integration.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.llms.anthropic.chat.handler import ModelResponseIterator
|
||||
from litellm.main import stream_chunk_builder
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
from litellm.types.responses.main import (
|
||||
OutputCodeInterpreterCall,
|
||||
OutputCodeInterpreterCallLog,
|
||||
)
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
|
||||
def _make_model_response(code_interpreter_results=None, provider_specific_fields=None):
|
||||
"""Helper to build a ModelResponse with provider_specific_fields on the message."""
|
||||
psf = provider_specific_fields or {}
|
||||
if code_interpreter_results is not None:
|
||||
psf["code_interpreter_results"] = code_interpreter_results
|
||||
msg = Message(content="test", provider_specific_fields=psf if psf else None)
|
||||
choice = Choices(index=0, message=msg, finish_reason="stop")
|
||||
resp = ModelResponse()
|
||||
resp.choices = [choice]
|
||||
return resp
|
||||
|
||||
|
||||
def test_extract_tool_result_output_items_from_pydantic_objects():
|
||||
"""Non-streaming path: code_interpreter_results are Pydantic OutputCodeInterpreterCall objects."""
|
||||
items = [
|
||||
OutputCodeInterpreterCall(
|
||||
type="code_interpreter_call",
|
||||
id="srvtoolu_01AAA",
|
||||
code="echo hello",
|
||||
container_id=None,
|
||||
status="completed",
|
||||
outputs=[OutputCodeInterpreterCallLog(type="logs", logs="hello\n")],
|
||||
),
|
||||
OutputCodeInterpreterCall(
|
||||
type="code_interpreter_call",
|
||||
id="srvtoolu_01BBB",
|
||||
code="echo world",
|
||||
container_id=None,
|
||||
status="completed",
|
||||
outputs=[OutputCodeInterpreterCallLog(type="logs", logs="world\n")],
|
||||
),
|
||||
]
|
||||
resp = _make_model_response(code_interpreter_results=items)
|
||||
result = LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(resp)
|
||||
assert len(result) == 2
|
||||
assert result[0].id == "srvtoolu_01AAA"
|
||||
assert result[1].id == "srvtoolu_01BBB"
|
||||
|
||||
|
||||
def test_extract_tool_result_output_items_from_dicts():
|
||||
"""Streaming path: after model_dump(), code_interpreter_results are plain dicts.
|
||||
_extract_tool_result_output_items reconstructs them as Pydantic objects."""
|
||||
items = [
|
||||
{
|
||||
"type": "code_interpreter_call",
|
||||
"id": "srvtoolu_01AAA",
|
||||
"code": "echo hello",
|
||||
"container_id": None,
|
||||
"status": "completed",
|
||||
"outputs": [{"type": "logs", "logs": "hello\n"}],
|
||||
},
|
||||
]
|
||||
resp = _make_model_response(code_interpreter_results=items)
|
||||
result = LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(resp)
|
||||
assert len(result) == 1
|
||||
assert isinstance(result[0], OutputCodeInterpreterCall)
|
||||
assert result[0].id == "srvtoolu_01AAA"
|
||||
|
||||
|
||||
def test_extract_tool_result_output_items_empty():
|
||||
"""No code_interpreter_results → empty list."""
|
||||
resp = _make_model_response()
|
||||
result = LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(resp)
|
||||
assert result == []
|
||||
|
||||
|
||||
def test_extract_tool_result_output_items_no_provider_specific_fields():
|
||||
"""Message with no provider_specific_fields → empty list."""
|
||||
msg = Message(content="test")
|
||||
choice = Choices(index=0, message=msg, finish_reason="stop")
|
||||
resp = ModelResponse()
|
||||
resp.choices = [choice]
|
||||
result = LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(resp)
|
||||
assert result == []
|
||||
|
||||
|
||||
def test_in_place_substitution_preserves_ordering():
|
||||
"""
|
||||
function_call items matching code_interpreter_results should be replaced
|
||||
in-place, preserving the original output ordering.
|
||||
|
||||
Simulates: [message, function_call(exec1), function_call(regular), function_call(exec2)]
|
||||
Expected: [message, code_interpreter_call(exec1), function_call(regular), code_interpreter_call(exec2)]
|
||||
"""
|
||||
code_results = [
|
||||
OutputCodeInterpreterCall(
|
||||
type="code_interpreter_call",
|
||||
id="srvtoolu_01AAA",
|
||||
code="echo first",
|
||||
container_id=None,
|
||||
status="completed",
|
||||
outputs=[OutputCodeInterpreterCallLog(type="logs", logs="first\n")],
|
||||
),
|
||||
OutputCodeInterpreterCall(
|
||||
type="code_interpreter_call",
|
||||
id="srvtoolu_01CCC",
|
||||
code="echo third",
|
||||
container_id=None,
|
||||
status="completed",
|
||||
outputs=[OutputCodeInterpreterCallLog(type="logs", logs="third\n")],
|
||||
),
|
||||
]
|
||||
resp = _make_model_response(code_interpreter_results=code_results)
|
||||
|
||||
# Build a mock responses_output list with interleaved items
|
||||
class MockItem:
|
||||
def __init__(self, type, call_id=None):
|
||||
self.type = type
|
||||
self.call_id = call_id
|
||||
|
||||
msg_item = MockItem(type="message")
|
||||
fc_exec1 = MockItem(type="function_call", call_id="srvtoolu_01AAA")
|
||||
fc_regular = MockItem(type="function_call", call_id="srvtoolu_01BBB")
|
||||
fc_exec2 = MockItem(type="function_call", call_id="srvtoolu_01CCC")
|
||||
|
||||
responses_output = [msg_item, fc_exec1, fc_regular, fc_exec2]
|
||||
|
||||
# Apply the same logic as _transform_chat_completion_choices_to_responses_output
|
||||
tool_result_items = (
|
||||
LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(resp)
|
||||
)
|
||||
if tool_result_items:
|
||||
result_by_id = {
|
||||
(item.get("id") if isinstance(item, dict) else item.id): item
|
||||
for item in tool_result_items
|
||||
}
|
||||
replaced_ids = set(result_by_id.keys())
|
||||
responses_output = [
|
||||
(
|
||||
result_by_id[getattr(item, "call_id", None)]
|
||||
if (
|
||||
getattr(item, "type", None) == "function_call"
|
||||
and getattr(item, "call_id", None) in replaced_ids
|
||||
)
|
||||
else item
|
||||
)
|
||||
for item in responses_output
|
||||
]
|
||||
|
||||
# Verify ordering: message, code_interpreter(AAA), function_call(BBB), code_interpreter(CCC)
|
||||
assert len(responses_output) == 4
|
||||
assert responses_output[0].type == "message"
|
||||
assert responses_output[1].type == "code_interpreter_call"
|
||||
assert responses_output[1].id == "srvtoolu_01AAA"
|
||||
assert responses_output[2].type == "function_call"
|
||||
assert responses_output[2].call_id == "srvtoolu_01BBB"
|
||||
assert responses_output[3].type == "code_interpreter_call"
|
||||
assert responses_output[3].id == "srvtoolu_01CCC"
|
||||
|
||||
|
||||
def test_end_to_end_streaming_chunks_to_code_interpreter_output():
|
||||
"""
|
||||
Mock end-to-end test: Anthropic SSE chunks → ModelResponseIterator →
|
||||
stream_chunk_builder → _extract_tool_result_output_items → final output
|
||||
with code_interpreter_call items replacing function_call items.
|
||||
|
||||
This exercises the full streaming data flow without a live server.
|
||||
"""
|
||||
# Realistic Anthropic streaming chunks for a single code execution
|
||||
raw_chunks = [
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": "msg_01XYZ",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [],
|
||||
"usage": {"input_tokens": 100, "output_tokens": 1},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 0,
|
||||
"content_block": {
|
||||
"type": "server_tool_use",
|
||||
"id": "srvtoolu_01AAA",
|
||||
"name": "bash_code_execution",
|
||||
"input": {},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {
|
||||
"type": "input_json_delta",
|
||||
"partial_json": '{"command": "echo e2e_test"}',
|
||||
},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": 1,
|
||||
"content_block": {
|
||||
"type": "bash_code_execution_tool_result",
|
||||
"tool_use_id": "srvtoolu_01AAA",
|
||||
"content": {
|
||||
"type": "bash_code_execution_result",
|
||||
"stdout": "e2e_test\n",
|
||||
"stderr": "",
|
||||
"return_code": 0,
|
||||
},
|
||||
},
|
||||
},
|
||||
{"type": "content_block_stop", "index": 1},
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn"},
|
||||
"usage": {"output_tokens": 50},
|
||||
},
|
||||
]
|
||||
|
||||
# Step 1: Parse chunks through ModelResponseIterator (Anthropic handler)
|
||||
iterator = ModelResponseIterator(None, sync_stream=True)
|
||||
parsed_chunks = []
|
||||
for chunk in raw_chunks:
|
||||
parsed = iterator.chunk_parser(chunk)
|
||||
d = parsed.model_dump()
|
||||
# In production, CustomStreamWrapper sets the model on each chunk;
|
||||
# stream_chunk_builder requires it.
|
||||
d["model"] = "claude-sonnet-4-20250514"
|
||||
parsed_chunks.append(d)
|
||||
|
||||
# Step 2: Assemble via stream_chunk_builder (simulates end-of-stream)
|
||||
assembled = stream_chunk_builder(chunks=parsed_chunks)
|
||||
assert assembled is not None
|
||||
|
||||
# Verify stream_chunk_builder picked up code_interpreter_results via last-value-wins
|
||||
psf = assembled.choices[0].message.provider_specific_fields
|
||||
assert psf is not None
|
||||
assert "code_interpreter_results" in psf
|
||||
code_results = psf["code_interpreter_results"]
|
||||
assert len(code_results) == 1
|
||||
# After model_dump + stream_chunk_builder, results are plain dicts
|
||||
assert code_results[0]["id"] == "srvtoolu_01AAA"
|
||||
assert code_results[0]["code"] == "echo e2e_test"
|
||||
|
||||
# Step 3: Extract via _extract_tool_result_output_items (Responses API layer)
|
||||
tool_result_items = (
|
||||
LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(assembled)
|
||||
)
|
||||
assert len(tool_result_items) == 1
|
||||
item = tool_result_items[0]
|
||||
# Items are reconstructed as Pydantic OutputCodeInterpreterCall objects
|
||||
assert isinstance(item, OutputCodeInterpreterCall)
|
||||
assert item.type == "code_interpreter_call"
|
||||
assert item.id == "srvtoolu_01AAA"
|
||||
assert item.code == "echo e2e_test"
|
||||
assert item.outputs[0].logs == "e2e_test\n"
|
||||
@@ -1666,3 +1666,141 @@ async def test_oauth_authorize_prefers_request_scope_over_server_config():
|
||||
redirect_url = response.headers["location"]
|
||||
assert "scope=custom_scope1+custom_scope2" in redirect_url or "scope=custom_scope1%20custom_scope2" in redirect_url
|
||||
assert "default_scope" not in redirect_url
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_endpoint_refresh_token_grant():
|
||||
"""Test that token endpoint supports refresh_token grant type."""
|
||||
try:
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
token_endpoint,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
except ImportError:
|
||||
pytest.skip("MCP discoverable endpoints not available")
|
||||
|
||||
# Clear registry
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
# Create mock OAuth2 server
|
||||
oauth2_server = MCPServer(
|
||||
server_id="google_mcp",
|
||||
name="google_mcp",
|
||||
server_name="google_mcp",
|
||||
alias="google_mcp",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="test_client_id",
|
||||
client_secret="test_secret",
|
||||
authorization_url="https://accounts.google.com/o/oauth2/v2/auth",
|
||||
token_url="https://oauth2.googleapis.com/token",
|
||||
scopes=["openid", "email"],
|
||||
)
|
||||
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "https://proxy.litellm.example/"
|
||||
mock_request.headers = {}
|
||||
|
||||
# Mock httpx client response with new tokens
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"access_token": "new_access_token",
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3599,
|
||||
"refresh_token": "new_refresh_token",
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_async_client = MagicMock()
|
||||
mock_async_client.post = AsyncMock(return_value=mock_response)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client"
|
||||
) as mock_get_client:
|
||||
mock_get_client.return_value = mock_async_client
|
||||
|
||||
response = await token_endpoint(
|
||||
request=mock_request,
|
||||
grant_type="refresh_token",
|
||||
code=None,
|
||||
redirect_uri=None,
|
||||
client_id="test_client_id",
|
||||
mcp_server_name="google_mcp",
|
||||
client_secret="test_secret",
|
||||
refresh_token="rt-test",
|
||||
scope="openid email",
|
||||
)
|
||||
|
||||
# Verify the POST was called with refresh_token grant data
|
||||
mock_async_client.post.assert_called_once()
|
||||
call_args = mock_async_client.post.call_args
|
||||
|
||||
assert call_args[1]["data"]["grant_type"] == "refresh_token"
|
||||
assert call_args[1]["data"]["refresh_token"] == "rt-test"
|
||||
assert call_args[1]["data"]["client_id"] == "test_client_id"
|
||||
assert call_args[1]["data"]["client_secret"] == "test_secret"
|
||||
assert call_args[1]["data"]["scope"] == "openid email"
|
||||
|
||||
# Verify response contains the new tokens
|
||||
import json
|
||||
|
||||
token_data = json.loads(response.body)
|
||||
assert token_data["access_token"] == "new_access_token"
|
||||
assert token_data["refresh_token"] == "new_refresh_token"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_token_endpoint_authorization_code_missing_code():
|
||||
"""Test that authorization_code grant rejects missing code param."""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
|
||||
exchange_token_with_server,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
from litellm.proxy._types import MCPTransport
|
||||
from litellm.types.mcp import MCPAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
except ImportError:
|
||||
pytest.skip("MCP discoverable endpoints not available")
|
||||
|
||||
global_mcp_server_manager.registry.clear()
|
||||
|
||||
server = MCPServer(
|
||||
server_id="test_server",
|
||||
name="test_server",
|
||||
server_name="test_server",
|
||||
alias="test_server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
client_id="cid",
|
||||
token_url="https://example.com/token",
|
||||
)
|
||||
global_mcp_server_manager.registry[server.server_id] = server
|
||||
|
||||
mock_request = MagicMock()
|
||||
mock_request.base_url = "https://proxy.example/"
|
||||
mock_request.headers = {}
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await exchange_token_with_server(
|
||||
request=mock_request,
|
||||
mcp_server=server,
|
||||
grant_type="authorization_code",
|
||||
code=None,
|
||||
redirect_uri="https://example.com/cb",
|
||||
client_id="cid",
|
||||
client_secret=None,
|
||||
code_verifier=None,
|
||||
)
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "code is required" in str(exc_info.value.detail)
|
||||
|
||||
@@ -1458,10 +1458,6 @@ async def test_unblock_key_supports_both_sk_and_hashed_tokens(monkeypatch):
|
||||
return_value=mock_key_record
|
||||
)
|
||||
|
||||
# Mock get_key_object and _cache_key_object functions
|
||||
mock_key_object = MagicMock()
|
||||
mock_key_object.blocked = True # Initially blocked
|
||||
|
||||
# Mock hash_token function
|
||||
def mock_hash_token(token):
|
||||
if token == "sk-test123456789":
|
||||
@@ -1482,19 +1478,12 @@ async def test_unblock_key_supports_both_sk_and_hashed_tokens(monkeypatch):
|
||||
) # Disable audit logs for simpler test
|
||||
|
||||
# Mock get_key_object and _cache_key_object
|
||||
async def mock_get_key_object(**kwargs):
|
||||
return mock_key_object
|
||||
|
||||
async def mock_cache_key_object(**kwargs):
|
||||
async def mock_delete_cache_key_object(**kwargs):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.get_key_object",
|
||||
mock_get_key_object,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._cache_key_object",
|
||||
mock_cache_key_object,
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
|
||||
mock_delete_cache_key_object,
|
||||
)
|
||||
|
||||
# Create mock request and user auth
|
||||
@@ -1519,11 +1508,9 @@ async def test_unblock_key_supports_both_sk_and_hashed_tokens(monkeypatch):
|
||||
)
|
||||
|
||||
assert result == mock_key_record
|
||||
assert mock_key_object.blocked == False # Should be updated to unblocked
|
||||
|
||||
# Reset mocks for second test
|
||||
mock_prisma_client.db.litellm_verificationtoken.update.reset_mock()
|
||||
mock_key_object.blocked = True # Reset to blocked state
|
||||
|
||||
# Test Case 2: Using already hashed token
|
||||
hashed_token_request = BlockKeyRequest(key=test_hashed_token)
|
||||
@@ -1541,7 +1528,6 @@ async def test_unblock_key_supports_both_sk_and_hashed_tokens(monkeypatch):
|
||||
)
|
||||
|
||||
assert result == mock_key_record
|
||||
assert mock_key_object.blocked == False # Should be updated to unblocked
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -1579,6 +1565,249 @@ async def test_unblock_key_invalid_key_format(monkeypatch):
|
||||
assert "Invalid key format" in str(exc_info.value.message)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_block_key_nonexistent_key_returns_404(monkeypatch):
|
||||
"""
|
||||
Test that block_key returns 404 (not misleading 401) when the key
|
||||
doesn't exist in the database, even when the caller is authenticated
|
||||
as a proxy admin.
|
||||
|
||||
Previously, block_key would call get_key_object() for cache refresh,
|
||||
which raised a 401 ProxyException with 'Authentication Error' — making
|
||||
it look like an auth failure when it was really a missing-key error.
|
||||
"""
|
||||
from litellm.proxy._types import BlockKeyRequest
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import block_key
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_user_api_key_cache = MagicMock()
|
||||
mock_proxy_logging_obj = MagicMock()
|
||||
|
||||
# find_unique returns None → key does not exist
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
|
||||
def mock_hash_token(token):
|
||||
return "abcd1234" * 8 # 64-char hex
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", mock_hash_token)
|
||||
monkeypatch.setattr("litellm.store_audit_logs", False)
|
||||
|
||||
mock_request = MagicMock()
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin_user"
|
||||
)
|
||||
|
||||
data = BlockKeyRequest(key="sk-does-not-exist-key")
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await block_key(
|
||||
data=data,
|
||||
http_request=mock_request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.code == "404"
|
||||
assert "not found" in str(exc_info.value.message).lower()
|
||||
# Must NOT contain "Authentication Error"
|
||||
assert "Authentication Error" not in str(exc_info.value.message)
|
||||
# update should never be called since the key doesn't exist
|
||||
mock_prisma_client.db.litellm_verificationtoken.update.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unblock_key_nonexistent_key_returns_404(monkeypatch):
|
||||
"""
|
||||
Test that unblock_key returns 404 (not misleading 401) when the key
|
||||
doesn't exist in the database.
|
||||
"""
|
||||
from litellm.proxy._types import BlockKeyRequest
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
unblock_key,
|
||||
)
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_user_api_key_cache = MagicMock()
|
||||
mock_proxy_logging_obj = MagicMock()
|
||||
|
||||
# find_unique returns None → key does not exist
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
|
||||
def mock_hash_token(token):
|
||||
return "abcd1234" * 8
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", mock_hash_token)
|
||||
monkeypatch.setattr("litellm.store_audit_logs", False)
|
||||
|
||||
mock_request = MagicMock()
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin_user"
|
||||
)
|
||||
|
||||
data = BlockKeyRequest(key="sk-does-not-exist-key")
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await unblock_key(
|
||||
data=data,
|
||||
http_request=mock_request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.code == "404"
|
||||
assert "not found" in str(exc_info.value.message).lower()
|
||||
assert "Authentication Error" not in str(exc_info.value.message)
|
||||
mock_prisma_client.db.litellm_verificationtoken.update.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_key_nonexistent_key_returns_404(monkeypatch):
|
||||
"""
|
||||
Test that update_key_fn returns 404 (not misleading 401) when the body
|
||||
key doesn't exist in the database, even when the caller is authenticated
|
||||
as a proxy admin via the Authorization header.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
update_key_fn,
|
||||
)
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_user_api_key_cache = MagicMock()
|
||||
mock_proxy_logging_obj = MagicMock()
|
||||
|
||||
# find_unique returns None → key does not exist
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
|
||||
|
||||
mock_request = MagicMock()
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin_user"
|
||||
)
|
||||
|
||||
data = UpdateKeyRequest(key="sk-does-not-exist-key")
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await update_key_fn(
|
||||
request=mock_request,
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.code == "404"
|
||||
assert "not found" in str(exc_info.value.message).lower()
|
||||
assert "Authentication Error" not in str(exc_info.value.message)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_block_key_existing_key_succeeds(monkeypatch):
|
||||
"""
|
||||
Test that block_key successfully blocks an existing key and
|
||||
invalidates the cache entry.
|
||||
"""
|
||||
from litellm.proxy._types import BlockKeyRequest
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import block_key
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_user_api_key_cache = MagicMock()
|
||||
mock_proxy_logging_obj = MagicMock()
|
||||
|
||||
test_hashed_token = "a1b2c3d4e5f6789012345678901234567890123456789012345678901234abcd"
|
||||
|
||||
mock_key_record = MagicMock()
|
||||
mock_key_record.token = test_hashed_token
|
||||
mock_key_record.blocked = False
|
||||
mock_key_record.model_dump_json.return_value = (
|
||||
f'{{"token": "{test_hashed_token}", "blocked": false}}'
|
||||
)
|
||||
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
|
||||
return_value=mock_key_record
|
||||
)
|
||||
mock_updated_record = MagicMock()
|
||||
mock_updated_record.token = test_hashed_token
|
||||
mock_updated_record.blocked = True
|
||||
mock_prisma_client.db.litellm_verificationtoken.update = AsyncMock(
|
||||
return_value=mock_updated_record
|
||||
)
|
||||
|
||||
def mock_hash_token(token):
|
||||
if token.startswith("sk-"):
|
||||
return test_hashed_token
|
||||
return token
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", mock_hash_token)
|
||||
monkeypatch.setattr("litellm.store_audit_logs", False)
|
||||
|
||||
# Mock _delete_cache_key_object
|
||||
async def mock_delete_cache_key_object(**kwargs):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
|
||||
mock_delete_cache_key_object,
|
||||
)
|
||||
|
||||
mock_request = MagicMock()
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin_user"
|
||||
)
|
||||
|
||||
data = BlockKeyRequest(key="sk-test123456789")
|
||||
|
||||
result = await block_key(
|
||||
data=data,
|
||||
http_request=mock_request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
# Verify the key was found and updated
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique.assert_called_once_with(
|
||||
where={"token": test_hashed_token}
|
||||
)
|
||||
mock_prisma_client.db.litellm_verificationtoken.update.assert_called_once_with(
|
||||
where={"token": test_hashed_token}, data={"blocked": True}
|
||||
)
|
||||
assert result == mock_updated_record
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_validate_key_team_change_with_member_permissions():
|
||||
"""
|
||||
@@ -4871,14 +5100,16 @@ async def test_validate_max_budget():
|
||||
async def test_get_and_validate_existing_key():
|
||||
"""
|
||||
Test _get_and_validate_existing_key helper function.
|
||||
|
||||
|
||||
Tests:
|
||||
1. Successfully retrieve existing key
|
||||
2. Key not found raises HTTPException
|
||||
2. Key not found raises ProxyException
|
||||
3. Database not connected raises HTTPException
|
||||
"""
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy._types import ProxyException
|
||||
|
||||
# Test Case 1: Successfully retrieve existing key
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_key = LiteLLM_VerificationToken(
|
||||
@@ -4887,39 +5118,49 @@ async def test_get_and_validate_existing_key():
|
||||
models=["gpt-4"],
|
||||
team_id=None,
|
||||
)
|
||||
mock_prisma_client.get_data = AsyncMock(return_value=mock_key)
|
||||
|
||||
result = await _get_and_validate_existing_key(
|
||||
token="test-key-123",
|
||||
prisma_client=mock_prisma_client,
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
|
||||
return_value=mock_key
|
||||
)
|
||||
|
||||
assert result == mock_key
|
||||
mock_prisma_client.get_data.assert_called_once_with(
|
||||
token="test-key-123",
|
||||
table_name="key",
|
||||
query_type="find_unique",
|
||||
)
|
||||
|
||||
# Test Case 2: Key not found raises HTTPException
|
||||
mock_prisma_client.get_data = AsyncMock(return_value=None)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _get_and_validate_existing_key(
|
||||
token="non-existent-key",
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed",
|
||||
return_value="hashed-test-key-123",
|
||||
):
|
||||
result = await _get_and_validate_existing_key(
|
||||
token="test-key-123",
|
||||
prisma_client=mock_prisma_client,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 404
|
||||
assert "Key not found" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
assert result == mock_key
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique.assert_called_once_with(
|
||||
where={"token": "hashed-test-key-123"}
|
||||
)
|
||||
|
||||
# Test Case 2: Key not found raises ProxyException
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed",
|
||||
return_value="hashed-non-existent-key",
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await _get_and_validate_existing_key(
|
||||
token="non-existent-key",
|
||||
prisma_client=mock_prisma_client,
|
||||
)
|
||||
|
||||
assert str(exc_info.value.code) == "404"
|
||||
assert "Key not found" in exc_info.value.message
|
||||
|
||||
# Test Case 3: Database not connected raises HTTPException
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await _get_and_validate_existing_key(
|
||||
token="test-key-123",
|
||||
prisma_client=None,
|
||||
)
|
||||
|
||||
|
||||
assert exc_info.value.status_code == 500
|
||||
assert "Database not connected" in str(exc_info.value.detail)
|
||||
|
||||
@@ -4960,75 +5201,82 @@ async def test_process_single_key_update():
|
||||
"tags": ["production"],
|
||||
}
|
||||
|
||||
mock_prisma_client.get_data = AsyncMock(return_value=existing_key)
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
|
||||
return_value=existing_key
|
||||
)
|
||||
mock_updated_key_obj = MagicMock()
|
||||
mock_updated_key_obj.model_dump.return_value = updated_key_data
|
||||
mock_prisma_client.update_data = AsyncMock(
|
||||
return_value={"data": mock_updated_key_obj}
|
||||
)
|
||||
|
||||
|
||||
# Mock prepare_key_update_data
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.prepare_key_update_data"
|
||||
) as mock_prepare:
|
||||
mock_prepare.return_value = {"max_budget": 100.0, "tags": ["production"]}
|
||||
|
||||
|
||||
# Mock TeamMemberPermissionChecks
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint"
|
||||
) as mock_permission_check:
|
||||
mock_permission_check.return_value = None
|
||||
|
||||
|
||||
# Mock _delete_cache_key_object
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object"
|
||||
) as mock_delete_cache:
|
||||
mock_delete_cache.return_value = None
|
||||
|
||||
|
||||
# Mock hash_token (imported from litellm.proxy._types)
|
||||
with patch(
|
||||
"litellm.proxy._types.hash_token"
|
||||
) as mock_hash:
|
||||
mock_hash.return_value = "hashed-test-key-123"
|
||||
|
||||
# Mock KeyManagementEventHooks
|
||||
|
||||
# Mock _hash_token_if_needed
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed",
|
||||
return_value="hashed-test-key-123",
|
||||
):
|
||||
# Create update request
|
||||
key_update_item = BulkUpdateKeyRequestItem(
|
||||
key="test-key-123",
|
||||
max_budget=100.0,
|
||||
tags=["production"],
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-admin",
|
||||
user_id="admin-user",
|
||||
)
|
||||
|
||||
# Call the function
|
||||
result = await _process_single_key_update(
|
||||
key_update_item=key_update_item,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=None,
|
||||
prisma_client=mock_prisma_client,
|
||||
user_api_key_cache=mock_user_api_key_cache,
|
||||
proxy_logging_obj=mock_proxy_logging_obj,
|
||||
llm_router=mock_llm_router,
|
||||
)
|
||||
|
||||
# Verify results
|
||||
assert result is not None
|
||||
assert "token" not in result # Token should be removed
|
||||
assert result.get("max_budget") == 100.0
|
||||
assert result.get("tags") == ["production"]
|
||||
|
||||
# Verify mocks were called
|
||||
mock_prisma_client.get_data.assert_called_once()
|
||||
mock_prisma_client.update_data.assert_called_once()
|
||||
mock_delete_cache.assert_called_once()
|
||||
# Mock KeyManagementEventHooks
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
|
||||
):
|
||||
# Create update request
|
||||
key_update_item = BulkUpdateKeyRequestItem(
|
||||
key="test-key-123",
|
||||
max_budget=100.0,
|
||||
tags=["production"],
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-admin",
|
||||
user_id="admin-user",
|
||||
)
|
||||
|
||||
# Call the function
|
||||
result = await _process_single_key_update(
|
||||
key_update_item=key_update_item,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=None,
|
||||
prisma_client=mock_prisma_client,
|
||||
user_api_key_cache=mock_user_api_key_cache,
|
||||
proxy_logging_obj=mock_proxy_logging_obj,
|
||||
llm_router=mock_llm_router,
|
||||
)
|
||||
|
||||
# Verify results
|
||||
assert result is not None
|
||||
assert "token" not in result # Token should be removed
|
||||
assert result.get("max_budget") == 100.0
|
||||
assert result.get("tags") == ["production"]
|
||||
|
||||
# Verify mocks were called
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique.assert_called_once()
|
||||
mock_prisma_client.update_data.assert_called_once()
|
||||
mock_delete_cache.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -5090,7 +5338,7 @@ async def test_bulk_update_keys_success(monkeypatch):
|
||||
"tags": ["staging"],
|
||||
}
|
||||
|
||||
mock_prisma_client.get_data = AsyncMock(
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
|
||||
side_effect=[existing_key_1, existing_key_2]
|
||||
)
|
||||
mock_updated_key_1_obj = MagicMock()
|
||||
@@ -5103,7 +5351,7 @@ async def test_bulk_update_keys_success(monkeypatch):
|
||||
{"data": mock_updated_key_2_obj},
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
# Patch dependencies
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.prisma_client", mock_prisma_client
|
||||
@@ -5115,7 +5363,7 @@ async def test_bulk_update_keys_success(monkeypatch):
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_llm_router)
|
||||
|
||||
|
||||
# Mock helper functions
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.prepare_key_update_data"
|
||||
@@ -5124,7 +5372,7 @@ async def test_bulk_update_keys_success(monkeypatch):
|
||||
{"max_budget": 100.0, "tags": ["production"]},
|
||||
{"max_budget": 200.0, "tags": ["staging"]},
|
||||
]
|
||||
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint"
|
||||
):
|
||||
@@ -5135,45 +5383,49 @@ async def test_bulk_update_keys_success(monkeypatch):
|
||||
"litellm.proxy._types.hash_token"
|
||||
) as mock_hash:
|
||||
mock_hash.side_effect = ["hashed-key-1", "hashed-key-2"]
|
||||
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed",
|
||||
side_effect=["hashed-key-1", "hashed-key-2"],
|
||||
):
|
||||
# Create request
|
||||
request_data = BulkUpdateKeyRequest(
|
||||
keys=[
|
||||
BulkUpdateKeyRequestItem(
|
||||
key="test-key-1",
|
||||
max_budget=100.0,
|
||||
tags=["production"],
|
||||
),
|
||||
BulkUpdateKeyRequestItem(
|
||||
key="test-key-2",
|
||||
max_budget=200.0,
|
||||
tags=["staging"],
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-admin",
|
||||
user_id="admin-user",
|
||||
)
|
||||
|
||||
# Call endpoint
|
||||
response = await bulk_update_keys(
|
||||
data=request_data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
# Verify response
|
||||
assert response.total_requested == 2
|
||||
assert len(response.successful_updates) == 2
|
||||
assert len(response.failed_updates) == 0
|
||||
assert response.successful_updates[0].key == "test-key-1"
|
||||
assert response.successful_updates[1].key == "test-key-2"
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
|
||||
):
|
||||
# Create request
|
||||
request_data = BulkUpdateKeyRequest(
|
||||
keys=[
|
||||
BulkUpdateKeyRequestItem(
|
||||
key="test-key-1",
|
||||
max_budget=100.0,
|
||||
tags=["production"],
|
||||
),
|
||||
BulkUpdateKeyRequestItem(
|
||||
key="test-key-2",
|
||||
max_budget=200.0,
|
||||
tags=["staging"],
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-admin",
|
||||
user_id="admin-user",
|
||||
)
|
||||
|
||||
# Call endpoint
|
||||
response = await bulk_update_keys(
|
||||
data=request_data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
# Verify response
|
||||
assert response.total_requested == 2
|
||||
assert len(response.successful_updates) == 2
|
||||
assert len(response.failed_updates) == 0
|
||||
assert response.successful_updates[0].key == "test-key-1"
|
||||
assert response.successful_updates[1].key == "test-key-2"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -5218,7 +5470,7 @@ async def test_bulk_update_keys_partial_failures(monkeypatch):
|
||||
}
|
||||
|
||||
# First key exists, second key doesn't exist
|
||||
mock_prisma_client.get_data = AsyncMock(
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
|
||||
side_effect=[existing_key_1, None] # Second key not found
|
||||
)
|
||||
mock_updated_key_1_obj = MagicMock()
|
||||
@@ -5226,7 +5478,9 @@ async def test_bulk_update_keys_partial_failures(monkeypatch):
|
||||
mock_prisma_client.update_data = AsyncMock(
|
||||
return_value={"data": mock_updated_key_1_obj}
|
||||
)
|
||||
|
||||
# Mock get_data for the error handler path (used to fetch key_info on failure)
|
||||
mock_prisma_client.get_data = AsyncMock(return_value=None)
|
||||
|
||||
# Patch dependencies
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.prisma_client", mock_prisma_client
|
||||
@@ -5238,13 +5492,13 @@ async def test_bulk_update_keys_partial_failures(monkeypatch):
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_llm_router)
|
||||
|
||||
|
||||
# Mock helper functions
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.prepare_key_update_data"
|
||||
) as mock_prepare:
|
||||
mock_prepare.return_value = {"max_budget": 100.0, "tags": ["production"]}
|
||||
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint"
|
||||
):
|
||||
@@ -5255,46 +5509,50 @@ async def test_bulk_update_keys_partial_failures(monkeypatch):
|
||||
"litellm.proxy._types.hash_token"
|
||||
) as mock_hash:
|
||||
mock_hash.return_value = "hashed-key-1"
|
||||
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed",
|
||||
side_effect=["hashed-key-1", "hashed-non-existent-key"],
|
||||
):
|
||||
# Create request with one valid and one invalid key
|
||||
request_data = BulkUpdateKeyRequest(
|
||||
keys=[
|
||||
BulkUpdateKeyRequestItem(
|
||||
key="test-key-1",
|
||||
max_budget=100.0,
|
||||
tags=["production"],
|
||||
),
|
||||
BulkUpdateKeyRequestItem(
|
||||
key="non-existent-key",
|
||||
max_budget=200.0,
|
||||
tags=["staging"],
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-admin",
|
||||
user_id="admin-user",
|
||||
)
|
||||
|
||||
# Call endpoint
|
||||
response = await bulk_update_keys(
|
||||
data=request_data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
# Verify response
|
||||
assert response.total_requested == 2
|
||||
assert len(response.successful_updates) == 1
|
||||
assert len(response.failed_updates) == 1
|
||||
assert response.successful_updates[0].key == "test-key-1"
|
||||
assert response.failed_updates[0].key == "non-existent-key"
|
||||
assert "Key not found" in response.failed_updates[0].failed_reason
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
|
||||
):
|
||||
# Create request with one valid and one invalid key
|
||||
request_data = BulkUpdateKeyRequest(
|
||||
keys=[
|
||||
BulkUpdateKeyRequestItem(
|
||||
key="test-key-1",
|
||||
max_budget=100.0,
|
||||
tags=["production"],
|
||||
),
|
||||
BulkUpdateKeyRequestItem(
|
||||
key="non-existent-key",
|
||||
max_budget=200.0,
|
||||
tags=["staging"],
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-admin",
|
||||
user_id="admin-user",
|
||||
)
|
||||
|
||||
# Call endpoint
|
||||
response = await bulk_update_keys(
|
||||
data=request_data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
|
||||
# Verify response
|
||||
assert response.total_requested == 2
|
||||
assert len(response.successful_updates) == 1
|
||||
assert len(response.failed_updates) == 1
|
||||
assert response.successful_updates[0].key == "test-key-1"
|
||||
assert response.failed_updates[0].key == "non-existent-key"
|
||||
assert "Key not found" in response.failed_updates[0].failed_reason
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
@@ -7379,19 +7637,12 @@ def _setup_block_unblock_mocks(monkeypatch, mock_key_team_id=None):
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", mock_hash_token)
|
||||
monkeypatch.setattr("litellm.store_audit_logs", False)
|
||||
|
||||
async def mock_get_key_object(**kwargs):
|
||||
return mock_key_object
|
||||
|
||||
async def mock_cache_key_object(**kwargs):
|
||||
async def mock_delete_cache_key_object(**kwargs):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.get_key_object",
|
||||
mock_get_key_object,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._cache_key_object",
|
||||
mock_cache_key_object,
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
|
||||
mock_delete_cache_key_object,
|
||||
)
|
||||
|
||||
return mock_prisma_client, test_hashed_token
|
||||
@@ -7638,16 +7889,9 @@ async def test_update_key_non_budget_fields_allowed_for_internal_user(monkeypatc
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", mock_hash_token)
|
||||
|
||||
async def mock_cache_key_object(**kwargs):
|
||||
pass
|
||||
|
||||
async def mock_delete_cache_key_object(**kwargs):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._cache_key_object",
|
||||
mock_cache_key_object,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
|
||||
mock_delete_cache_key_object,
|
||||
|
||||
@@ -1519,6 +1519,8 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
client_id="client",
|
||||
client_secret="secret",
|
||||
code_verifier="verifier",
|
||||
refresh_token=None,
|
||||
scope=None,
|
||||
)
|
||||
|
||||
assert result is exchange_response
|
||||
@@ -1532,6 +1534,56 @@ class TestTemporaryMCPSessionEndpoints:
|
||||
client_id="client",
|
||||
client_secret="secret",
|
||||
code_verifier="verifier",
|
||||
refresh_token=None,
|
||||
scope=None,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_token_proxies_refresh_token_grant(self):
|
||||
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
|
||||
mcp_token,
|
||||
)
|
||||
|
||||
request = MagicMock()
|
||||
server = generate_mock_mcp_server_config_record(server_id="server-1")
|
||||
exchange_response = {"access_token": "new-token", "refresh_token": "new-rt"}
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
|
||||
return_value=server,
|
||||
) as get_server,
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
|
||||
AsyncMock(return_value=exchange_response),
|
||||
) as exchange_mock,
|
||||
):
|
||||
result = await mcp_token(
|
||||
request=request,
|
||||
server_id="server-1",
|
||||
grant_type="refresh_token",
|
||||
code=None,
|
||||
redirect_uri=None,
|
||||
client_id="client",
|
||||
client_secret="secret",
|
||||
code_verifier=None,
|
||||
refresh_token="rt-123",
|
||||
scope=None,
|
||||
)
|
||||
|
||||
assert result is exchange_response
|
||||
get_server.assert_called_once_with("server-1")
|
||||
exchange_mock.assert_awaited_once_with(
|
||||
request=request,
|
||||
mcp_server=server,
|
||||
grant_type="refresh_token",
|
||||
code=None,
|
||||
redirect_uri=None,
|
||||
client_id="client",
|
||||
client_secret="secret",
|
||||
code_verifier=None,
|
||||
refresh_token="rt-123",
|
||||
scope=None,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@@ -280,6 +280,47 @@ class TestProxyInitializationHelpers:
|
||||
assert result.exit_code == 0, f"exit_code={result.exit_code}, output={result.output}"
|
||||
mock_uvicorn_run.assert_called_once()
|
||||
|
||||
@patch("uvicorn.run")
|
||||
@patch("atexit.register")
|
||||
@patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database")
|
||||
@patch("litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=False)
|
||||
def test_proxy_default_api_version_uses_azure_default(
|
||||
self, mock_should_update, mock_setup_db, mock_atexit_register, mock_uvicorn_run
|
||||
):
|
||||
"""Proxy default api_version should match litellm.AZURE_DEFAULT_API_VERSION for consistency."""
|
||||
from click.testing import CliRunner
|
||||
|
||||
import litellm
|
||||
from litellm.proxy.proxy_cli import run_server
|
||||
|
||||
runner = CliRunner()
|
||||
mock_proxy_module = MagicMock(
|
||||
app=MagicMock(),
|
||||
ProxyConfig=MagicMock(),
|
||||
KeyManagementSettings=MagicMock(),
|
||||
save_worker_config=MagicMock(),
|
||||
)
|
||||
clean_env = {k: v for k, v in os.environ.items() if k not in ("DATABASE_URL", "DIRECT_URL")}
|
||||
with patch.dict(os.environ, clean_env, clear=True), patch.dict(
|
||||
"sys.modules",
|
||||
{
|
||||
"proxy_server": mock_proxy_module,
|
||||
"litellm.proxy.proxy_server": mock_proxy_module,
|
||||
},
|
||||
), patch(
|
||||
"litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args"
|
||||
) as mock_get_args:
|
||||
mock_get_args.return_value = {
|
||||
"app": "litellm.proxy.proxy_server:app",
|
||||
"host": "localhost",
|
||||
"port": 8000,
|
||||
}
|
||||
result = runner.invoke(run_server, ["--local", "--skip_server_startup"])
|
||||
assert result.exit_code == 0, f"exit_code={result.exit_code}, output={result.output}"
|
||||
mock_proxy_module.save_worker_config.assert_called_once()
|
||||
call_kwargs = mock_proxy_module.save_worker_config.call_args[1]
|
||||
assert call_kwargs["api_version"] == litellm.AZURE_DEFAULT_API_VERSION
|
||||
|
||||
@patch("uvicorn.run")
|
||||
@patch("builtins.print")
|
||||
def test_keepalive_timeout_flag(self, mock_print, mock_uvicorn_run):
|
||||
|
||||
@@ -0,0 +1,146 @@
|
||||
import { render, screen } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { vi } from "vitest";
|
||||
import { flexRender, getCoreRowModel, useReactTable } from "@tanstack/react-table";
|
||||
import { getAgentHubTableColumns, AgentHubData } from "./AgentHubTableColumns";
|
||||
|
||||
const mockAgent: AgentHubData = {
|
||||
agent_id: "agent-1",
|
||||
protocolVersion: "1.0",
|
||||
name: "Test Agent",
|
||||
description: "A test agent for unit testing",
|
||||
url: "https://agent.example.com",
|
||||
version: "2.0",
|
||||
capabilities: { streaming: true, caching: false },
|
||||
defaultInputModes: ["text"],
|
||||
defaultOutputModes: ["text", "image"],
|
||||
skills: [
|
||||
{ id: "s1", name: "Skill One", description: "First skill" },
|
||||
{ id: "s2", name: "Skill Two", description: "Second skill" },
|
||||
{ id: "s3", name: "Skill Three", description: "Third skill" },
|
||||
],
|
||||
is_public: true,
|
||||
};
|
||||
|
||||
function TestTable({
|
||||
data,
|
||||
publicPage = false,
|
||||
showModal = vi.fn(),
|
||||
copyToClipboard = vi.fn(),
|
||||
}: {
|
||||
data: AgentHubData[];
|
||||
publicPage?: boolean;
|
||||
showModal?: ReturnType<typeof vi.fn>;
|
||||
copyToClipboard?: ReturnType<typeof vi.fn>;
|
||||
}) {
|
||||
const columns = getAgentHubTableColumns(showModal, copyToClipboard, publicPage);
|
||||
const table = useReactTable({ data, columns, getCoreRowModel: getCoreRowModel() });
|
||||
|
||||
return (
|
||||
<table>
|
||||
<thead>
|
||||
{table.getHeaderGroups().map((hg) => (
|
||||
<tr key={hg.id}>
|
||||
{hg.headers.map((h) => (
|
||||
<th key={h.id}>{flexRender(h.column.columnDef.header, h.getContext())}</th>
|
||||
))}
|
||||
</tr>
|
||||
))}
|
||||
</thead>
|
||||
<tbody>
|
||||
{table.getRowModel().rows.map((row) => (
|
||||
<tr key={row.id}>
|
||||
{row.getVisibleCells().map((cell) => (
|
||||
<td key={cell.id}>{flexRender(cell.column.columnDef.cell, cell.getContext())}</td>
|
||||
))}
|
||||
</tr>
|
||||
))}
|
||||
</tbody>
|
||||
</table>
|
||||
);
|
||||
}
|
||||
|
||||
describe("AgentHubTableColumns", () => {
|
||||
it("should render", () => {
|
||||
render(<TestTable data={[mockAgent]} />);
|
||||
expect(screen.getByText("Test Agent")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display the agent description", () => {
|
||||
render(<TestTable data={[mockAgent]} />);
|
||||
// Description appears in both the description column and the mobile view within agent name column
|
||||
expect(screen.getAllByText("A test agent for unit testing").length).toBeGreaterThanOrEqual(1);
|
||||
});
|
||||
|
||||
it("should display the version with a 'v' prefix", () => {
|
||||
render(<TestTable data={[mockAgent]} />);
|
||||
expect(screen.getByText("v2.0")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display the protocol version", () => {
|
||||
render(<TestTable data={[mockAgent]} />);
|
||||
expect(screen.getByText("1.0")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show skill count with correct pluralization", () => {
|
||||
render(<TestTable data={[mockAgent]} />);
|
||||
expect(screen.getByText("3 skills")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show first two skills and '+1' for overflow", () => {
|
||||
render(<TestTable data={[mockAgent]} />);
|
||||
expect(screen.getByText("Skill One")).toBeInTheDocument();
|
||||
expect(screen.getByText("Skill Two")).toBeInTheDocument();
|
||||
expect(screen.getByText("+1")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show only true capabilities as badges", () => {
|
||||
render(<TestTable data={[mockAgent]} />);
|
||||
expect(screen.getByText("streaming")).toBeInTheDocument();
|
||||
expect(screen.queryByText("caching")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display I/O modes", () => {
|
||||
render(<TestTable data={[mockAgent]} />);
|
||||
// "In:" and "Out:" are in <span> children; getByText with exact:false
|
||||
// matches against the element's full textContent across child nodes
|
||||
expect(screen.getByText((_, el) =>
|
||||
el?.tagName === "P" && el.textContent === "In: text"
|
||||
)).toBeInTheDocument();
|
||||
expect(screen.getByText((_, el) =>
|
||||
el?.tagName === "P" && el.textContent === "Out: text, image"
|
||||
)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display 'Yes' badge for public agents", () => {
|
||||
render(<TestTable data={[mockAgent]} />);
|
||||
expect(screen.getByText("Yes")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display 'No' badge for non-public agents", () => {
|
||||
const privateAgent = { ...mockAgent, is_public: false };
|
||||
render(<TestTable data={[privateAgent]} />);
|
||||
expect(screen.getByText("No")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display a Details button", () => {
|
||||
render(<TestTable data={[mockAgent]} />);
|
||||
expect(screen.getByRole("button", { name: /details|info/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show '-' when agent has no capabilities", () => {
|
||||
const noCapAgent = { ...mockAgent, capabilities: {} };
|
||||
render(<TestTable data={[noCapAgent]} />);
|
||||
// The dash is rendered in the capabilities column
|
||||
expect(screen.getByText("-")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show singular 'skill' for one skill", () => {
|
||||
const oneSkillAgent = {
|
||||
...mockAgent,
|
||||
skills: [{ id: "s1", name: "Only Skill", description: "One" }],
|
||||
};
|
||||
render(<TestTable data={[oneSkillAgent]} />);
|
||||
expect(screen.getByText("1 skill")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
@@ -194,7 +194,6 @@ export const getAgentHubTableColumns = (
|
||||
return publicA - publicB;
|
||||
},
|
||||
cell: ({ row }) => {
|
||||
console.log(`CHECKPOINT 1: ${JSON.stringify(row.original)}`);
|
||||
const agent = row.original;
|
||||
|
||||
return agent.is_public === true ? (
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
import { renderWithProviders, screen } from "../../../tests/test-utils";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { vi } from "vitest";
|
||||
import UsageExportHeader from "./UsageExportHeader";
|
||||
import type { EntitySpendData } from "./types";
|
||||
|
||||
vi.mock("./EntityUsageExportModal", () => ({
|
||||
default: ({ isOpen, onClose }: { isOpen: boolean; onClose: () => void }) =>
|
||||
isOpen ? (
|
||||
<div data-testid="export-modal">
|
||||
<button onClick={onClose}>Close</button>
|
||||
</div>
|
||||
) : null,
|
||||
}));
|
||||
|
||||
const defaultProps = {
|
||||
dateValue: { from: new Date("2025-01-01"), to: new Date("2025-01-31") },
|
||||
entityType: "team" as const,
|
||||
spendData: {
|
||||
results: [],
|
||||
metadata: {
|
||||
total_spend: 0,
|
||||
total_api_requests: 0,
|
||||
total_successful_requests: 0,
|
||||
total_failed_requests: 0,
|
||||
total_tokens: 0,
|
||||
},
|
||||
} satisfies EntitySpendData,
|
||||
};
|
||||
|
||||
describe("UsageExportHeader", () => {
|
||||
it("should render", () => {
|
||||
renderWithProviders(<UsageExportHeader {...defaultProps} />);
|
||||
expect(screen.getByRole("button", { name: /export data/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should open the export modal when the export button is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<UsageExportHeader {...defaultProps} />);
|
||||
await user.click(screen.getByRole("button", { name: /export data/i }));
|
||||
expect(screen.getByTestId("export-modal")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should close the export modal when onClose is called", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<UsageExportHeader {...defaultProps} />);
|
||||
await user.click(screen.getByRole("button", { name: /export data/i }));
|
||||
await user.click(screen.getByRole("button", { name: /close/i }));
|
||||
expect(screen.queryByTestId("export-modal")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should not show filter dropdown when showFilters is false", () => {
|
||||
renderWithProviders(<UsageExportHeader {...defaultProps} showFilters={false} />);
|
||||
expect(screen.queryByText(/filter/i)).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show filter dropdown when showFilters is true and options provided", () => {
|
||||
renderWithProviders(
|
||||
<UsageExportHeader
|
||||
{...defaultProps}
|
||||
showFilters
|
||||
filterLabel="Team"
|
||||
filterPlaceholder="Select teams"
|
||||
filterOptions={[
|
||||
{ label: "Team A", value: "team-a" },
|
||||
{ label: "Team B", value: "team-b" },
|
||||
]}
|
||||
onFiltersChange={vi.fn()}
|
||||
/>,
|
||||
);
|
||||
expect(screen.getByText("Team")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,98 @@
|
||||
import { render, screen, act } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { vi } from "vitest";
|
||||
import { GuardrailConfig } from "./GuardrailConfig";
|
||||
|
||||
describe("GuardrailConfig", () => {
|
||||
const defaultProps = {
|
||||
guardrailName: "Content Safety",
|
||||
guardrailType: "Content Safety",
|
||||
provider: "bedrock",
|
||||
};
|
||||
|
||||
afterEach(() => {
|
||||
vi.useRealTimers();
|
||||
});
|
||||
|
||||
it("should render", () => {
|
||||
render(<GuardrailConfig {...defaultProps} />);
|
||||
expect(screen.getByText("Parameters")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display the guardrail name in the parameters description", () => {
|
||||
render(<GuardrailConfig {...defaultProps} />);
|
||||
expect(screen.getByText(/Configure Content Safety behavior/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
// Note: Version history entries are hardcoded placeholders in the component.
|
||||
// These assertions will need updating when wired to real API data.
|
||||
it("should show version history when 'View history' is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<GuardrailConfig {...defaultProps} />);
|
||||
await user.click(screen.getByRole("button", { name: /view history/i }));
|
||||
expect(screen.getByText("Initial configuration")).toBeInTheDocument();
|
||||
expect(screen.getByText("Added custom categories list")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should toggle version history text between View/Hide", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<GuardrailConfig {...defaultProps} />);
|
||||
const button = screen.getByRole("button", { name: /view history/i });
|
||||
await user.click(button);
|
||||
expect(screen.getByRole("button", { name: /hide history/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show custom code textarea when custom code override is toggled on", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<GuardrailConfig {...defaultProps} />);
|
||||
// Walk up from "Custom Code Override" heading to find the enclosing section,
|
||||
// then locate the switch within it
|
||||
const heading = screen.getByText("Custom Code Override");
|
||||
let container = heading.parentElement;
|
||||
let customCodeSwitch: Element | null = null;
|
||||
while (container && !customCodeSwitch) {
|
||||
customCodeSwitch = container.querySelector('[role="switch"]');
|
||||
container = container.parentElement;
|
||||
}
|
||||
if (!customCodeSwitch) {
|
||||
throw new Error("Could not find the Custom Code Override switch via DOM traversal");
|
||||
}
|
||||
await user.click(customCodeSwitch);
|
||||
expect(screen.getByPlaceholderText(/async def evaluate/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should hide custom code textarea when custom code override is off", () => {
|
||||
render(<GuardrailConfig {...defaultProps} />);
|
||||
// There's an input for categories, but no textarea
|
||||
expect(screen.queryByPlaceholderText(/async def evaluate/)).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show the re-run button in idle state", () => {
|
||||
render(<GuardrailConfig {...defaultProps} />);
|
||||
expect(screen.getByRole("button", { name: /re-run on failing logs/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show loading state when re-run is clicked", async () => {
|
||||
vi.useFakeTimers({ shouldAdvanceTime: true });
|
||||
const user = userEvent.setup({ advanceTimers: vi.advanceTimersByTime });
|
||||
render(<GuardrailConfig {...defaultProps} />);
|
||||
await user.click(screen.getByRole("button", { name: /re-run on failing logs/i }));
|
||||
expect(screen.getByText(/Running on 10 samples/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show success message after re-run completes", async () => {
|
||||
vi.useFakeTimers({ shouldAdvanceTime: true });
|
||||
const user = userEvent.setup({ advanceTimers: vi.advanceTimersByTime });
|
||||
render(<GuardrailConfig {...defaultProps} />);
|
||||
await user.click(screen.getByRole("button", { name: /re-run on failing logs/i }));
|
||||
await act(async () => { vi.advanceTimersByTime(2500); });
|
||||
expect(screen.getByText(/7\/10 would now pass/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display the Revert and Save buttons", () => {
|
||||
render(<GuardrailConfig {...defaultProps} />);
|
||||
expect(screen.getByRole("button", { name: /revert/i })).toBeInTheDocument();
|
||||
// The component's hardcoded default version is "v3", so Save shows "v4"
|
||||
expect(screen.getByRole("button", { name: /save as v\d+/i })).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
@@ -138,4 +138,18 @@ describe("DocsMenu", () => {
|
||||
await user.click(button);
|
||||
expect(button).toHaveAttribute("aria-expanded", "true");
|
||||
});
|
||||
|
||||
it("should close menu when clicking outside", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<div>
|
||||
<DocsMenu items={items} />
|
||||
<button>Outside</button>
|
||||
</div>,
|
||||
);
|
||||
await user.click(screen.getByRole("button", { name: /docs/i }));
|
||||
expect(screen.getByText("Custom pricing")).toBeInTheDocument();
|
||||
await user.click(screen.getByRole("button", { name: /outside/i }));
|
||||
expect(screen.queryByText("Custom pricing")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
@@ -26,6 +26,10 @@ const PERMISSION_OPTIONS = [
|
||||
"/key/unblock",
|
||||
"/key/bulk_update",
|
||||
"/key/{key_id}/reset_spend",
|
||||
"/key/info",
|
||||
"/key/list",
|
||||
"/key/aliases",
|
||||
"/team/daily/activity",
|
||||
];
|
||||
|
||||
interface SettingRowProps {
|
||||
|
||||
@@ -13,6 +13,7 @@ import {
|
||||
CreditCardOutlined,
|
||||
DatabaseOutlined,
|
||||
ExperimentOutlined,
|
||||
ExportOutlined,
|
||||
FileTextOutlined,
|
||||
FolderOutlined,
|
||||
KeyOutlined,
|
||||
@@ -400,7 +401,7 @@ const Sidebar: React.FC<SidebarProps> = ({ setPage, defaultSelectedKey, collapse
|
||||
onClick={(e) => e.stopPropagation()}
|
||||
style={{ color: "inherit", textDecoration: "none" }}
|
||||
>
|
||||
{label}
|
||||
{label} <ExportOutlined style={{ fontSize: 10, marginLeft: 4 }} />
|
||||
</a>
|
||||
);
|
||||
}
|
||||
|
||||
@@ -0,0 +1,296 @@
|
||||
import { render, screen } from "@testing-library/react";
|
||||
import { describe, it, expect, vi } from "vitest";
|
||||
import ChatMessageBubble from "./ChatMessageBubble";
|
||||
import { EndpointType } from "./mode_endpoint_mapping";
|
||||
import { MessageType } from "./types";
|
||||
|
||||
// Mock child components to isolate bubble rendering logic
|
||||
vi.mock("react-markdown", () => ({
|
||||
default: ({ children }: { children: string }) => <div data-testid="react-markdown">{children}</div>,
|
||||
}));
|
||||
|
||||
vi.mock("react-syntax-highlighter", () => ({
|
||||
Prism: ({ children }: { children: string }) => <pre data-testid="syntax-highlighter">{children}</pre>,
|
||||
}));
|
||||
|
||||
vi.mock("react-syntax-highlighter/dist/esm/styles/prism", () => ({
|
||||
coy: {},
|
||||
}));
|
||||
|
||||
vi.mock("./ReasoningContent", () => ({
|
||||
default: ({ reasoningContent }: { reasoningContent: string }) => (
|
||||
<div data-testid="reasoning-content">{reasoningContent}</div>
|
||||
),
|
||||
}));
|
||||
|
||||
vi.mock("./MCPEventsDisplay", () => ({
|
||||
default: ({ events }: { events: unknown[] }) => (
|
||||
<div data-testid="mcp-events-display">{events.length} events</div>
|
||||
),
|
||||
}));
|
||||
|
||||
vi.mock("./SearchResultsDisplay", () => ({
|
||||
SearchResultsDisplay: ({ searchResults }: { searchResults: unknown[] }) => (
|
||||
<div data-testid="search-results-display">{searchResults.length} results</div>
|
||||
),
|
||||
}));
|
||||
|
||||
vi.mock("./ResponseMetrics", () => ({
|
||||
default: ({ timeToFirstToken }: { timeToFirstToken?: number }) => (
|
||||
<div data-testid="response-metrics">TTFT: {timeToFirstToken}</div>
|
||||
),
|
||||
}));
|
||||
|
||||
vi.mock("./A2AMetrics", () => ({
|
||||
default: ({ a2aMetadata }: { a2aMetadata: unknown }) => (
|
||||
<div data-testid="a2a-metrics">A2A</div>
|
||||
),
|
||||
}));
|
||||
|
||||
vi.mock("./CodeInterpreterOutput", () => ({
|
||||
default: ({ code }: { code: string }) => <div data-testid="code-interpreter-output">{code}</div>,
|
||||
}));
|
||||
|
||||
vi.mock("./AudioRenderer", () => ({
|
||||
default: ({ message }: { message: MessageType }) => (
|
||||
<div data-testid="audio-renderer">{typeof message.content === "string" ? message.content : ""}</div>
|
||||
),
|
||||
}));
|
||||
|
||||
vi.mock("./ResponsesImageRenderer", () => ({
|
||||
default: () => <div data-testid="responses-image-renderer" />,
|
||||
}));
|
||||
|
||||
vi.mock("./ChatImageRenderer", () => ({
|
||||
default: () => <div data-testid="chat-image-renderer" />,
|
||||
}));
|
||||
|
||||
const defaultProps = {
|
||||
isLastMessage: false,
|
||||
endpointType: EndpointType.CHAT,
|
||||
mcpEvents: [],
|
||||
codeInterpreterResult: null,
|
||||
accessToken: "test-token",
|
||||
};
|
||||
|
||||
describe("ChatMessageBubble", () => {
|
||||
it("should render a user message with right-aligned text", () => {
|
||||
render(
|
||||
<ChatMessageBubble
|
||||
{...defaultProps}
|
||||
message={{ role: "user", content: "Hello" }}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.getByText("user")).toBeInTheDocument();
|
||||
expect(screen.getByText("Hello")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render an assistant message with left-aligned text", () => {
|
||||
render(
|
||||
<ChatMessageBubble
|
||||
{...defaultProps}
|
||||
message={{ role: "assistant", content: "Hi there" }}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.getByText("assistant")).toBeInTheDocument();
|
||||
expect(screen.getByText("Hi there")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show model badge for assistant messages when model is provided", () => {
|
||||
render(
|
||||
<ChatMessageBubble
|
||||
{...defaultProps}
|
||||
message={{ role: "assistant", content: "Reply", model: "gpt-4" }}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.getByText("gpt-4")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should not show model badge for user messages even when model is set", () => {
|
||||
render(
|
||||
<ChatMessageBubble
|
||||
{...defaultProps}
|
||||
message={{ role: "user", content: "Hello", model: "gpt-4" }}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.queryByText("gpt-4")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render markdown content via ReactMarkdown", () => {
|
||||
render(
|
||||
<ChatMessageBubble
|
||||
{...defaultProps}
|
||||
message={{ role: "assistant", content: "**bold text**" }}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.getByTestId("react-markdown")).toHaveTextContent("**bold text**");
|
||||
});
|
||||
|
||||
it("should render an image when isImage is true", () => {
|
||||
render(
|
||||
<ChatMessageBubble
|
||||
{...defaultProps}
|
||||
message={{ role: "assistant", content: "https://example.com/img.png", isImage: true }}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.getByAltText("Generated image")).toHaveAttribute("src", "https://example.com/img.png");
|
||||
});
|
||||
|
||||
it("should render AudioRenderer when isAudio is true", () => {
|
||||
render(
|
||||
<ChatMessageBubble
|
||||
{...defaultProps}
|
||||
message={{ role: "assistant", content: "audio-url", isAudio: true }}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.getByTestId("audio-renderer")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show ReasoningContent when reasoningContent is present", () => {
|
||||
render(
|
||||
<ChatMessageBubble
|
||||
{...defaultProps}
|
||||
message={{ role: "assistant", content: "answer", reasoningContent: "thinking..." }}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.getByTestId("reasoning-content")).toHaveTextContent("thinking...");
|
||||
});
|
||||
|
||||
it("should show MCP events on the last assistant message for RESPONSES endpoint", () => {
|
||||
const mcpEvents = [{ type: "tool_call", item_id: "1" }];
|
||||
|
||||
render(
|
||||
<ChatMessageBubble
|
||||
{...defaultProps}
|
||||
isLastMessage={true}
|
||||
endpointType={EndpointType.RESPONSES}
|
||||
mcpEvents={mcpEvents as any}
|
||||
message={{ role: "assistant", content: "response" }}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.getByTestId("mcp-events-display")).toHaveTextContent("1 events");
|
||||
});
|
||||
|
||||
it("should show MCP events on the last assistant message for CHAT endpoint", () => {
|
||||
const mcpEvents = [{ type: "tool_call", item_id: "1" }];
|
||||
|
||||
render(
|
||||
<ChatMessageBubble
|
||||
{...defaultProps}
|
||||
isLastMessage={true}
|
||||
endpointType={EndpointType.CHAT}
|
||||
mcpEvents={mcpEvents as any}
|
||||
message={{ role: "assistant", content: "response" }}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.getByTestId("mcp-events-display")).toHaveTextContent("1 events");
|
||||
});
|
||||
|
||||
it("should not show MCP events when isLastMessage is false", () => {
|
||||
const mcpEvents = [{ type: "tool_call", item_id: "1" }];
|
||||
|
||||
render(
|
||||
<ChatMessageBubble
|
||||
{...defaultProps}
|
||||
isLastMessage={false}
|
||||
endpointType={EndpointType.RESPONSES}
|
||||
mcpEvents={mcpEvents as any}
|
||||
message={{ role: "assistant", content: "response" }}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.queryByTestId("mcp-events-display")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show SearchResultsDisplay when searchResults are present", () => {
|
||||
render(
|
||||
<ChatMessageBubble
|
||||
{...defaultProps}
|
||||
message={{
|
||||
role: "assistant",
|
||||
content: "found results",
|
||||
searchResults: [{ object: "search", search_query: "q", data: [] }],
|
||||
}}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.getByTestId("search-results-display")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show ResponseMetrics when usage data is present and no a2aMetadata", () => {
|
||||
render(
|
||||
<ChatMessageBubble
|
||||
{...defaultProps}
|
||||
message={{
|
||||
role: "assistant",
|
||||
content: "response",
|
||||
timeToFirstToken: 150,
|
||||
usage: { completionTokens: 10, promptTokens: 5, totalTokens: 15 },
|
||||
}}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.getByTestId("response-metrics")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show A2AMetrics when a2aMetadata is present instead of ResponseMetrics", () => {
|
||||
render(
|
||||
<ChatMessageBubble
|
||||
{...defaultProps}
|
||||
message={{
|
||||
role: "assistant",
|
||||
content: "agent response",
|
||||
timeToFirstToken: 100,
|
||||
a2aMetadata: { taskId: "task-1", status: { state: "completed" } },
|
||||
}}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.getByTestId("a2a-metrics")).toBeInTheDocument();
|
||||
expect(screen.queryByTestId("response-metrics")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show CodeInterpreterOutput on the last assistant message for RESPONSES endpoint", () => {
|
||||
render(
|
||||
<ChatMessageBubble
|
||||
{...defaultProps}
|
||||
isLastMessage={true}
|
||||
endpointType={EndpointType.RESPONSES}
|
||||
codeInterpreterResult={{
|
||||
code: "print('hello')",
|
||||
containerId: "container-1",
|
||||
annotations: [],
|
||||
}}
|
||||
message={{ role: "assistant", content: "result" }}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.getByTestId("code-interpreter-output")).toHaveTextContent("print('hello')");
|
||||
});
|
||||
|
||||
it("should render generated image from chat completions via message.image", () => {
|
||||
render(
|
||||
<ChatMessageBubble
|
||||
{...defaultProps}
|
||||
message={{
|
||||
role: "assistant",
|
||||
content: "Here is your image",
|
||||
image: { url: "https://example.com/generated.png", detail: "auto" },
|
||||
}}
|
||||
/>,
|
||||
);
|
||||
|
||||
const images = screen.getAllByAltText("Generated image");
|
||||
expect(images.some((img) => img.getAttribute("src") === "https://example.com/generated.png")).toBe(true);
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,214 @@
|
||||
import { RobotOutlined, UserOutlined } from "@ant-design/icons";
|
||||
import React from "react";
|
||||
import ReactMarkdown from "react-markdown";
|
||||
import { Prism as SyntaxHighlighter } from "react-syntax-highlighter";
|
||||
import { coy } from "react-syntax-highlighter/dist/esm/styles/prism";
|
||||
import { CodeInterpreterResult } from "../llm_calls/code_interpreter_handler";
|
||||
import A2AMetrics from "./A2AMetrics";
|
||||
import AudioRenderer from "./AudioRenderer";
|
||||
import ChatImageRenderer from "./ChatImageRenderer";
|
||||
import CodeInterpreterOutput from "./CodeInterpreterOutput";
|
||||
import { EndpointType } from "./mode_endpoint_mapping";
|
||||
import MCPEventsDisplay from "./MCPEventsDisplay";
|
||||
import type { MCPEvent } from "../../mcp_tools/types";
|
||||
import ReasoningContent from "./ReasoningContent";
|
||||
import ResponseMetrics from "./ResponseMetrics";
|
||||
import ResponsesImageRenderer from "./ResponsesImageRenderer";
|
||||
import { SearchResultsDisplay } from "./SearchResultsDisplay";
|
||||
import { MessageType } from "./types";
|
||||
|
||||
interface ChatMessageBubbleProps {
|
||||
message: MessageType;
|
||||
/** Whether this is the last message in the chat history. */
|
||||
isLastMessage: boolean;
|
||||
endpointType: EndpointType;
|
||||
/** MCP events to display on the last assistant message. */
|
||||
mcpEvents: MCPEvent[];
|
||||
/** Code interpreter result to display on the last assistant message. */
|
||||
codeInterpreterResult: CodeInterpreterResult | null;
|
||||
/** API key used to fetch code interpreter file downloads. */
|
||||
accessToken: string;
|
||||
}
|
||||
|
||||
function ChatMessageBubble({
|
||||
message,
|
||||
isLastMessage,
|
||||
endpointType,
|
||||
mcpEvents,
|
||||
codeInterpreterResult,
|
||||
accessToken,
|
||||
}: ChatMessageBubbleProps) {
|
||||
const isUser = message.role === "user";
|
||||
|
||||
return (
|
||||
<div className={`mb-4 ${isUser ? "text-right" : "text-left"}`}>
|
||||
<div
|
||||
className="inline-block max-w-[80%] rounded-lg shadow-sm p-3.5 px-4"
|
||||
style={{
|
||||
backgroundColor: isUser ? "#f0f8ff" : "#ffffff",
|
||||
border: isUser ? "1px solid #e6f0fa" : "1px solid #f0f0f0",
|
||||
textAlign: "left",
|
||||
}}
|
||||
>
|
||||
{/* Header: role icon + name + model badge */}
|
||||
<div className="flex items-center gap-2 mb-1.5">
|
||||
<div
|
||||
className="flex items-center justify-center w-6 h-6 rounded-full mr-1"
|
||||
style={{
|
||||
backgroundColor: isUser ? "#e6f0fa" : "#f5f5f5",
|
||||
}}
|
||||
>
|
||||
{isUser ? (
|
||||
<UserOutlined style={{ fontSize: "12px", color: "#2563eb" }} />
|
||||
) : (
|
||||
<RobotOutlined style={{ fontSize: "12px", color: "#4b5563" }} />
|
||||
)}
|
||||
</div>
|
||||
<strong className="text-sm capitalize">{message.role}</strong>
|
||||
{message.role === "assistant" && message.model && (
|
||||
<span className="text-xs px-2 py-0.5 rounded bg-gray-100 text-gray-600 font-normal">
|
||||
{message.model}
|
||||
</span>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* Reasoning content (chain-of-thought) */}
|
||||
{message.reasoningContent && <ReasoningContent reasoningContent={message.reasoningContent} />}
|
||||
|
||||
{/* MCP events at the start of the last assistant message */}
|
||||
{message.role === "assistant" &&
|
||||
isLastMessage &&
|
||||
mcpEvents.length > 0 &&
|
||||
(endpointType === EndpointType.RESPONSES || endpointType === EndpointType.CHAT) && (
|
||||
<div className="mb-3">
|
||||
<MCPEventsDisplay events={mcpEvents} />
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Search results */}
|
||||
{message.role === "assistant" && message.searchResults && (
|
||||
<SearchResultsDisplay searchResults={message.searchResults} />
|
||||
)}
|
||||
|
||||
{/* Code Interpreter output for the last assistant message */}
|
||||
{message.role === "assistant" &&
|
||||
isLastMessage &&
|
||||
codeInterpreterResult &&
|
||||
endpointType === EndpointType.RESPONSES && (
|
||||
<CodeInterpreterOutput
|
||||
code={codeInterpreterResult.code}
|
||||
containerId={codeInterpreterResult.containerId}
|
||||
annotations={codeInterpreterResult.annotations}
|
||||
accessToken={accessToken}
|
||||
/>
|
||||
)}
|
||||
|
||||
{/* Message body */}
|
||||
<div
|
||||
className="whitespace-pre-wrap break-words max-w-full message-content"
|
||||
style={{
|
||||
wordWrap: "break-word",
|
||||
overflowWrap: "break-word",
|
||||
wordBreak: "break-word",
|
||||
hyphens: "auto",
|
||||
}}
|
||||
>
|
||||
{message.isImage ? (
|
||||
<img
|
||||
src={typeof message.content === "string" ? message.content : ""}
|
||||
alt="Generated image"
|
||||
className="max-w-full rounded-md border border-gray-200 shadow-sm"
|
||||
style={{ maxHeight: "500px" }}
|
||||
/>
|
||||
) : message.isAudio ? (
|
||||
<AudioRenderer message={message} />
|
||||
) : (
|
||||
<>
|
||||
{/* Attached image for user messages based on endpoint */}
|
||||
{endpointType === EndpointType.RESPONSES && <ResponsesImageRenderer message={message} />}
|
||||
{endpointType === EndpointType.CHAT && <ChatImageRenderer message={message} />}
|
||||
|
||||
<ReactMarkdown
|
||||
components={{
|
||||
code({
|
||||
node,
|
||||
inline,
|
||||
className,
|
||||
children,
|
||||
...props
|
||||
}: React.ComponentPropsWithoutRef<"code"> & {
|
||||
inline?: boolean;
|
||||
node?: unknown;
|
||||
}) {
|
||||
const match = /language-(\w+)/.exec(className || "");
|
||||
return !inline && match ? (
|
||||
<SyntaxHighlighter
|
||||
style={coy as any}
|
||||
language={match[1]}
|
||||
PreTag="div"
|
||||
className="rounded-md my-2"
|
||||
wrapLines={true}
|
||||
wrapLongLines={true}
|
||||
{...props}
|
||||
>
|
||||
{String(children).replace(/\n$/, "")}
|
||||
</SyntaxHighlighter>
|
||||
) : (
|
||||
<code
|
||||
className={`${className} px-1.5 py-0.5 rounded bg-gray-100 text-sm font-mono`}
|
||||
style={{ wordBreak: "break-word" }}
|
||||
{...props}
|
||||
>
|
||||
{children}
|
||||
</code>
|
||||
);
|
||||
},
|
||||
pre: ({ node, ...props }) => (
|
||||
<pre style={{ overflowX: "auto", maxWidth: "100%" }} {...props} />
|
||||
),
|
||||
}}
|
||||
>
|
||||
{typeof message.content === "string" ? message.content : ""}
|
||||
</ReactMarkdown>
|
||||
|
||||
{/* Generated image from chat completions */}
|
||||
{message.image && (
|
||||
<div className="mt-3">
|
||||
<img
|
||||
src={message.image.url}
|
||||
alt="Generated image"
|
||||
className="max-w-full rounded-md border border-gray-200 shadow-sm"
|
||||
style={{ maxHeight: "500px" }}
|
||||
/>
|
||||
</div>
|
||||
)}
|
||||
</>
|
||||
)}
|
||||
|
||||
{/* Response metrics */}
|
||||
{message.role === "assistant" &&
|
||||
(message.timeToFirstToken || message.totalLatency || message.usage) &&
|
||||
!message.a2aMetadata && (
|
||||
<ResponseMetrics
|
||||
timeToFirstToken={message.timeToFirstToken}
|
||||
totalLatency={message.totalLatency}
|
||||
usage={message.usage}
|
||||
toolName={message.toolName}
|
||||
/>
|
||||
)}
|
||||
|
||||
{/* A2A Metrics */}
|
||||
{message.role === "assistant" && message.a2aMetadata && (
|
||||
<A2AMetrics
|
||||
a2aMetadata={message.a2aMetadata}
|
||||
timeToFirstToken={message.timeToFirstToken}
|
||||
totalLatency={message.totalLatency}
|
||||
/>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
export default ChatMessageBubble;
|
||||
@@ -63,6 +63,7 @@ import EndpointSelector from "./EndpointSelector";
|
||||
import FilePreviewCard from "./FilePreviewCard";
|
||||
import MCPEventsDisplay from "./MCPEventsDisplay";
|
||||
import type { MCPEvent } from "../../mcp_tools/types";
|
||||
import ChatMessageBubble from "./ChatMessageBubble";
|
||||
import { EndpointType, getEndpointType } from "./mode_endpoint_mapping";
|
||||
import ReasoningContent from "./ReasoningContent";
|
||||
import ResponseMetrics, { TokenUsage } from "./ResponseMetrics";
|
||||
@@ -1932,168 +1933,14 @@ const ChatUI: React.FC<ChatUIProps> = ({
|
||||
|
||||
{chatHistory.map((message, index) => (
|
||||
<div key={index}>
|
||||
<div className={`mb-4 ${message.role === "user" ? "text-right" : "text-left"}`}>
|
||||
<div
|
||||
className="inline-block max-w-[80%] rounded-lg shadow-sm p-3.5 px-4"
|
||||
style={{
|
||||
backgroundColor: message.role === "user" ? "#f0f8ff" : "#ffffff",
|
||||
border: message.role === "user" ? "1px solid #e6f0fa" : "1px solid #f0f0f0",
|
||||
textAlign: "left",
|
||||
}}
|
||||
>
|
||||
<div className="flex items-center gap-2 mb-1.5">
|
||||
<div
|
||||
className="flex items-center justify-center w-6 h-6 rounded-full mr-1"
|
||||
style={{
|
||||
backgroundColor: message.role === "user" ? "#e6f0fa" : "#f5f5f5",
|
||||
}}
|
||||
>
|
||||
{message.role === "user" ? (
|
||||
<UserOutlined style={{ fontSize: "12px", color: "#2563eb" }} />
|
||||
) : (
|
||||
<RobotOutlined style={{ fontSize: "12px", color: "#4b5563" }} />
|
||||
)}
|
||||
</div>
|
||||
<strong className="text-sm capitalize">{message.role}</strong>
|
||||
{message.role === "assistant" && message.model && (
|
||||
<span className="text-xs px-2 py-0.5 rounded bg-gray-100 text-gray-600 font-normal">
|
||||
{message.model}
|
||||
</span>
|
||||
)}
|
||||
</div>
|
||||
{message.reasoningContent && <ReasoningContent reasoningContent={message.reasoningContent} />}
|
||||
|
||||
{/* Show MCP events at the start of assistant messages */}
|
||||
{message.role === "assistant" &&
|
||||
index === chatHistory.length - 1 &&
|
||||
mcpEvents.length > 0 &&
|
||||
(endpointType === EndpointType.RESPONSES || endpointType === EndpointType.CHAT) && (
|
||||
<div className="mb-3">
|
||||
<MCPEventsDisplay events={mcpEvents} />
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Show search results at the start of assistant messages */}
|
||||
{message.role === "assistant" && message.searchResults && (
|
||||
<SearchResultsDisplay searchResults={message.searchResults} />
|
||||
)}
|
||||
|
||||
{/* Show Code Interpreter output for the last assistant message */}
|
||||
{message.role === "assistant" &&
|
||||
index === chatHistory.length - 1 &&
|
||||
codeInterpreter.result &&
|
||||
endpointType === EndpointType.RESPONSES && (
|
||||
<CodeInterpreterOutput
|
||||
code={codeInterpreter.result.code}
|
||||
containerId={codeInterpreter.result.containerId}
|
||||
annotations={codeInterpreter.result.annotations}
|
||||
accessToken={apiKeySource === "session" ? accessToken || "" : apiKey}
|
||||
/>
|
||||
)}
|
||||
|
||||
<div
|
||||
className="whitespace-pre-wrap break-words max-w-full message-content"
|
||||
style={{
|
||||
wordWrap: "break-word",
|
||||
overflowWrap: "break-word",
|
||||
wordBreak: "break-word",
|
||||
hyphens: "auto",
|
||||
}}
|
||||
>
|
||||
{message.isImage ? (
|
||||
<img
|
||||
src={typeof message.content === "string" ? message.content : ""}
|
||||
alt="Generated image"
|
||||
className="max-w-full rounded-md border border-gray-200 shadow-sm"
|
||||
style={{ maxHeight: "500px" }}
|
||||
/>
|
||||
) : message.isAudio ? (
|
||||
<AudioRenderer message={message} />
|
||||
) : (
|
||||
<>
|
||||
{/* Show attached image for user messages based on current endpoint */}
|
||||
{endpointType === EndpointType.RESPONSES && <ResponsesImageRenderer message={message} />}
|
||||
{endpointType === EndpointType.CHAT && <ChatImageRenderer message={message} />}
|
||||
|
||||
<ReactMarkdown
|
||||
components={{
|
||||
code({
|
||||
node,
|
||||
inline,
|
||||
className,
|
||||
children,
|
||||
...props
|
||||
}: React.ComponentPropsWithoutRef<"code"> & {
|
||||
inline?: boolean;
|
||||
node?: any;
|
||||
}) {
|
||||
const match = /language-(\w+)/.exec(className || "");
|
||||
return !inline && match ? (
|
||||
<SyntaxHighlighter
|
||||
style={coy as any}
|
||||
language={match[1]}
|
||||
PreTag="div"
|
||||
className="rounded-md my-2"
|
||||
wrapLines={true}
|
||||
wrapLongLines={true}
|
||||
{...props}
|
||||
>
|
||||
{String(children).replace(/\n$/, "")}
|
||||
</SyntaxHighlighter>
|
||||
) : (
|
||||
<code
|
||||
className={`${className} px-1.5 py-0.5 rounded bg-gray-100 text-sm font-mono`}
|
||||
style={{ wordBreak: "break-word" }}
|
||||
{...props}
|
||||
>
|
||||
{children}
|
||||
</code>
|
||||
);
|
||||
},
|
||||
pre: ({ node, ...props }) => (
|
||||
<pre style={{ overflowX: "auto", maxWidth: "100%" }} {...props} />
|
||||
),
|
||||
}}
|
||||
>
|
||||
{typeof message.content === "string" ? message.content : ""}
|
||||
</ReactMarkdown>
|
||||
|
||||
{/* Show generated image from chat completions */}
|
||||
{message.image && (
|
||||
<div className="mt-3">
|
||||
<img
|
||||
src={message.image.url}
|
||||
alt="Generated image"
|
||||
className="max-w-full rounded-md border border-gray-200 shadow-sm"
|
||||
style={{ maxHeight: "500px" }}
|
||||
/>
|
||||
</div>
|
||||
)}
|
||||
</>
|
||||
)}
|
||||
|
||||
{message.role === "assistant" &&
|
||||
(message.timeToFirstToken || message.totalLatency || message.usage) &&
|
||||
!message.a2aMetadata && (
|
||||
<ResponseMetrics
|
||||
timeToFirstToken={message.timeToFirstToken}
|
||||
totalLatency={message.totalLatency}
|
||||
usage={message.usage}
|
||||
toolName={message.toolName}
|
||||
/>
|
||||
)}
|
||||
|
||||
{/* A2A Metrics - show for A2A agent responses */}
|
||||
{message.role === "assistant" && message.a2aMetadata && (
|
||||
<A2AMetrics
|
||||
a2aMetadata={message.a2aMetadata}
|
||||
timeToFirstToken={message.timeToFirstToken}
|
||||
totalLatency={message.totalLatency}
|
||||
/>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
<ChatMessageBubble
|
||||
message={message}
|
||||
isLastMessage={index === chatHistory.length - 1}
|
||||
endpointType={endpointType as EndpointType}
|
||||
mcpEvents={mcpEvents}
|
||||
codeInterpreterResult={codeInterpreter.result}
|
||||
accessToken={apiKeySource === "session" ? accessToken || "" : apiKey}
|
||||
/>
|
||||
</div>
|
||||
))}
|
||||
|
||||
|
||||
+33
@@ -151,6 +151,39 @@ describe("GuardrailViewer", () => {
|
||||
expect(screen.queryByText(/Raw Bedrock Guardrail Response/)).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("renders without crashing when guardrail_mode is null", () => {
|
||||
const data = makeGuardrailInformation({ guardrail_mode: null });
|
||||
renderWithProviders(<GuardrailViewer data={data} />);
|
||||
|
||||
expect(screen.getByText("Guardrails & Policy Compliance")).toBeInTheDocument();
|
||||
// Null mode should display as dash
|
||||
expect(screen.getByText("—")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("renders without crashing when guardrail_mode is an object", () => {
|
||||
const data = makeGuardrailInformation({
|
||||
guardrail_mode: { default: "pre_call", tags: {} },
|
||||
});
|
||||
renderWithProviders(<GuardrailViewer data={data} />);
|
||||
|
||||
expect(screen.getByText("Guardrails & Policy Compliance")).toBeInTheDocument();
|
||||
expect(screen.getByText("PRE-CALL")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("renders without crashing when guardrail_mode is an array and shows in both timeline buckets", () => {
|
||||
const data = makeGuardrailInformation({
|
||||
guardrail_mode: ["pre_call", "post_call"],
|
||||
});
|
||||
renderWithProviders(<GuardrailViewer data={data} />);
|
||||
|
||||
expect(screen.getByText("Guardrails & Policy Compliance")).toBeInTheDocument();
|
||||
// Mode badge shows first element formatted
|
||||
expect(screen.getByText("PRE-CALL")).toBeInTheDocument();
|
||||
// Entry should appear in both pre-call and post-call timeline sections
|
||||
expect(screen.getByText(/Pre-call guardrail:/)).toBeInTheDocument();
|
||||
expect(screen.getByText(/Post-call guardrail:/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("integration: renders with real Bedrock details without mocks", async () => {
|
||||
const user = userEvent.setup();
|
||||
const data = makeGuardrailInformation({
|
||||
|
||||
@@ -40,7 +40,7 @@ interface GuardrailInformation {
|
||||
duration: number;
|
||||
end_time: number;
|
||||
start_time: number;
|
||||
guardrail_mode: string;
|
||||
guardrail_mode: string | string[] | Record<string, unknown> | null;
|
||||
guardrail_name: string;
|
||||
guardrail_status: string;
|
||||
guardrail_response: GuardrailEntity[] | BedrockGuardrailResponse | any;
|
||||
@@ -77,9 +77,50 @@ const PROVIDERS_WITH_CUSTOM_RENDERERS = new Set([
|
||||
"litellm_content_filter",
|
||||
]);
|
||||
|
||||
const formatMode = (mode: unknown): string => {
|
||||
if (mode == null || mode === "") return "—";
|
||||
const s = typeof mode === "string" ? mode : String(mode);
|
||||
/**
|
||||
* Extracts a plain string from guardrail_mode for display purposes.
|
||||
* Returns the first mode when multiple are present.
|
||||
*/
|
||||
const resolveMode = (mode: GuardrailInformation["guardrail_mode"]): string | null => {
|
||||
if (mode == null) return null;
|
||||
if (typeof mode === "string") return mode;
|
||||
if (Array.isArray(mode)) {
|
||||
const first = mode[0];
|
||||
return typeof first === "string" ? first : null;
|
||||
}
|
||||
if (typeof mode === "object" && "default" in mode) {
|
||||
const def = mode.default;
|
||||
if (typeof def === "string") return def;
|
||||
if (Array.isArray(def)) {
|
||||
const first = def[0];
|
||||
return typeof first === "string" ? first : null;
|
||||
}
|
||||
}
|
||||
return null;
|
||||
};
|
||||
|
||||
/**
|
||||
* Checks whether guardrail_mode includes the given target stage.
|
||||
* Handles arrays (multi-stage guardrails) by checking all elements.
|
||||
*/
|
||||
const modeMatches = (
|
||||
mode: GuardrailInformation["guardrail_mode"],
|
||||
target: string,
|
||||
): boolean => {
|
||||
if (mode == null) return false;
|
||||
if (typeof mode === "string") return mode === target;
|
||||
if (Array.isArray(mode)) return mode.includes(target);
|
||||
if (typeof mode === "object" && "default" in mode) {
|
||||
const def = mode.default;
|
||||
if (typeof def === "string") return def === target;
|
||||
if (Array.isArray(def)) return def.some((x) => typeof x === "string" && x === target);
|
||||
}
|
||||
return false;
|
||||
};
|
||||
|
||||
const formatMode = (mode: GuardrailInformation["guardrail_mode"]): string => {
|
||||
const s = resolveMode(mode);
|
||||
if (s == null || s === "") return "—";
|
||||
return s.replace(/_/g, "-").toUpperCase();
|
||||
};
|
||||
|
||||
@@ -301,10 +342,13 @@ const RequestLifecycle = ({ entries }: { entries: GuardrailInformation[] }) => {
|
||||
// Request received
|
||||
items.push({ type: "request", label: "Request received", offsetMs: 0 });
|
||||
|
||||
// Pre-call guardrails
|
||||
const preCalls = sorted.filter((e) => e.guardrail_mode === "pre_call");
|
||||
const postCalls = sorted.filter((e) => e.guardrail_mode === "post_call" || e.guardrail_mode === "logging_only");
|
||||
const duringCalls = sorted.filter((e) => e.guardrail_mode === "during_call");
|
||||
// Pre-call guardrails — use modeMatches so array modes (e.g. ["pre_call", "post_call"])
|
||||
// place the entry in every matching bucket.
|
||||
const preCalls = sorted.filter((e) => modeMatches(e.guardrail_mode, "pre_call"));
|
||||
const postCalls = sorted.filter(
|
||||
(e) => modeMatches(e.guardrail_mode, "post_call") || modeMatches(e.guardrail_mode, "logging_only"),
|
||||
);
|
||||
const duringCalls = sorted.filter((e) => modeMatches(e.guardrail_mode, "during_call"));
|
||||
|
||||
for (const e of preCalls) {
|
||||
const offsetMs = Math.round((e.end_time - baseTime) * 1000);
|
||||
|
||||
@@ -23,7 +23,7 @@ export interface GuardrailInformation {
|
||||
duration: number;
|
||||
end_time: number;
|
||||
start_time: number;
|
||||
guardrail_mode: string;
|
||||
guardrail_mode: string | string[] | Record<string, unknown> | null;
|
||||
guardrail_name: string;
|
||||
guardrail_status: string;
|
||||
guardrail_response: GuardrailEntity[] | BedrockGuardrailResponse;
|
||||
|
||||
Reference in New Issue
Block a user