Merge pull request #23818 from BerriAI/litellm_oss_staging_03_17_2026

fix(fireworks): skip #transform=inline for base64 data URLs (#23729)
This commit is contained in:
Krish Dholakia 2026-03-21 14:54:39 -07:00 • committed by GitHub
commit f911d8d865
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
41 changed files with 2768 additions and 265 deletions

View file

@ -0,0 +1,139 @@
# Akto
## Overview
[Akto](https://www.akto.io/) provides API security guardrails and data ingestion for LLM traffic.
Akto now uses a **two-entry guardrail pattern** in LiteLLM:
- `akto-validate` (`pre_call`) for request validation
- `akto-ingest` (`post_call`) for request/response ingestion
There is no `on_flagged` setting anymore.
Use these as two separate guardrails in `config.yaml`:
- `guardrail_name: "akto-validate"`
- `guardrail_name: "akto-ingest"`
## 1. Get Your Akto Credentials
Set up the Akto Guardrail API Service and grab:
- `AKTO_GUARDRAIL_API_BASE` — your Guardrail API Base URL
- `AKTO_API_KEY` — your API key
## 2. Configure in `config.yaml`
### Block + Ingest (recommended)
Use both entries below. This gives you:
- pre-call block decision
- post-call ingestion for allowed traffic
Keep these as two separate entries (`akto-validate` and `akto-ingest`).
```yaml
guardrails:
- guardrail_name: "akto-validate"
litellm_params:
guardrail: akto
mode: pre_call
akto_base_url: os.environ/AKTO_GUARDRAIL_API_BASE
akto_api_key: os.environ/AKTO_API_KEY
default_on: true
unreachable_fallback: fail_closed # optional: fail_open | fail_closed (default: fail_closed)
guardrail_timeout: 5 # optional, default: 5
akto_account_id: "1000000" # optional, env fallback: AKTO_ACCOUNT_ID
akto_vxlan_id: "0" # optional, env fallback: AKTO_VXLAN_ID
- guardrail_name: "akto-ingest"
litellm_params:
guardrail: akto
mode: post_call
akto_base_url: os.environ/AKTO_GUARDRAIL_API_BASE
akto_api_key: os.environ/AKTO_API_KEY
default_on: true
```
### Monitor-only mode
If you only want logging/ingestion and no blocking, keep only `akto-ingest`.
```yaml
guardrails:
- guardrail_name: "akto-ingest"
litellm_params:
guardrail: akto
mode: post_call
akto_base_url: os.environ/AKTO_GUARDRAIL_API_BASE
akto_api_key: os.environ/AKTO_API_KEY
default_on: true
```
## 3. Test It
```shell
curl -i http://localhost:4000/v1/chat/completions \
-H "Content-Type: application/json" \
-H "Authorization: Bearer <your litellm key>" \
-d '{
"model": "gpt-3.5-turbo",
"messages": [
{"role": "user", "content": "Hello, how are you?"}
]
}'
```
If a request gets blocked:
```json
{
"error": {
"message": "Prompt injection detected",
"type": "None",
"param": "None",
"code": "403"
}
}
```
## 4. How It Works
**Block + Ingest mode:**
```
Request → LiteLLM → Akto guardrail check
→ Allowed → forward to LLM → ingest response
→ Blocked → ingest blocked marker → 403 error
```
**Monitor-only mode:**
```
Request → LiteLLM → forward to LLM → get response
→ Send to Akto (guardrails + ingest) → log only
```
## 5. Event behavior
| Entry | LiteLLM hook | Akto call behavior |
|------|---|---|
| `akto-validate` | `pre_call` | Awaited call with `guardrails=true`, `ingest_data=false` |
| `akto-ingest` | `post_call` | Fire-and-forget call with `guardrails=true`, `ingest_data=true` |
When blocked in `pre_call`, LiteLLM sends one fire-and-forget ingest payload with blocked metadata and returns `403`.
## 6. Parameters
| Parameter | Env Variable | Default | Description |
|-----------|-------------|---------|-------------|
| `akto_base_url` | `AKTO_GUARDRAIL_API_BASE` | *required* | Akto Guardrail API Base URL |
| `akto_api_key` | `AKTO_API_KEY` | *required* | API key (sent as `Authorization` header) |
| `akto_account_id` | `AKTO_ACCOUNT_ID` | `1000000` | Akto account id included in payload |
| `akto_vxlan_id` | `AKTO_VXLAN_ID` | `0` | Akto vxlan id included in payload |
| `unreachable_fallback` | — | `fail_closed` | `fail_open` or `fail_closed` |
| `guardrail_timeout` | — | `5` | Timeout in seconds |
| `default_on` | — | `true` (recommended) | Enables the guardrail entry by default |
## 7. Error Handling
| Scenario | `fail_closed` (default) | `fail_open` |
|----------|------------------------|-------------|
| Akto unreachable | ❌ Blocked (503) | ✅ Passes through |
| Akto returns error | ❌ Blocked (503) | ✅ Passes through |
| Guardrail says no | ❌ Blocked (403) | ❌ Blocked (403) |

View file

@ -594,9 +594,26 @@ Expected Response
:::tip gpt-5.4: reasoning_effort + function tools
LiteLLM drops `reasoning_effort` from `gpt-5.4` requests to `litellm.completion()` that include tools, since that combination is supported in the Responses API.
When `gpt-5.4+` requests to `litellm.completion()` include both `reasoning_effort` and `tools`, LiteLLM **automatically routes** the request through the Responses API bridge. This works for both **OpenAI** (`openai/gpt-5.4`) and **Azure** (`azure/gpt-5.4`) providers — no extra configuration needed.
If you need reasoning **and** tools together, use `openai/responses/gpt-5.4` to route through the Responses API instead. See [Responses API Bridge](/docs/providers/openai#openai-chat-completion-to-responses-api-bridge) for details.
You can also route explicitly via `openai/responses/gpt-5.4` or `azure/responses/gpt-5.4`. See [Responses API Bridge](/docs/providers/openai#openai-chat-completion-to-responses-api-bridge) for details.
**Azure custom deployment names:** Auto-routing relies on the deployment name matching the `gpt-5.4*` pattern. If you use a custom deployment name (e.g. `"my-reasoning-model"`), enable routing via:
**SDK:**
```python
litellm.completion(model="azure/responses/my-reasoning-model", ...)
```
**Proxy config:**
```yaml
model_list:
- model_name: my-reasoning-model
litellm_params:
model: azure/my-reasoning-model
model_info:
mode: responses
```
:::

View file

@ -222,8 +222,12 @@ def _get_redis_client_logic(**env_overrides):
"REDIS_CLUSTER_NODES"
)
# If startup_nodes resolved to None (not set by kwarg or env), remove the key
# entirely so callers can rely on key presence as a reliable cluster-mode signal.
if _startup_nodes is not None and isinstance(_startup_nodes, str):
redis_kwargs["startup_nodes"] = json.loads(_startup_nodes)
elif _startup_nodes is None:
redis_kwargs.pop("startup_nodes", None)
_sentinel_nodes: Optional[Union[str, list]] = redis_kwargs.get("sentinel_nodes", None) or get_secret( # type: ignore
"REDIS_SENTINEL_NODES"
@ -273,10 +277,14 @@ def _get_redis_client_logic(**env_overrides):
redis_kwargs["ssl_ca_certs"] = _gcp_ssl_ca_certs
if "url" in redis_kwargs and redis_kwargs["url"] is not None:
redis_kwargs.pop("host", None)
redis_kwargs.pop("port", None)
redis_kwargs.pop("db", None)
redis_kwargs.pop("password", None)
# Only strip host/port/db/password when not routing to a cluster.
# When startup_nodes is also present the cluster path takes priority and
# needs the password for authentication.
if not redis_kwargs.get("startup_nodes"):
redis_kwargs.pop("host", None)
redis_kwargs.pop("port", None)
redis_kwargs.pop("db", None)
redis_kwargs.pop("password", None)
elif "startup_nodes" in redis_kwargs and redis_kwargs["startup_nodes"] is not None:
pass
elif (
@ -368,6 +376,10 @@ def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis:
def get_redis_client(**env_overrides):
redis_kwargs = _get_redis_client_logic(**env_overrides)
if "startup_nodes" in redis_kwargs:
return init_redis_cluster(redis_kwargs)
if "url" in redis_kwargs and redis_kwargs["url"] is not None:
args = _get_redis_url_kwargs()
url_kwargs = {}
@ -377,9 +389,6 @@ def get_redis_client(**env_overrides):
return redis.Redis.from_url(**url_kwargs)
if "startup_nodes" in redis_kwargs or get_secret("REDIS_CLUSTER_NODES") is not None: # type: ignore
return init_redis_cluster(redis_kwargs)
# Check for Redis Sentinel
if "sentinel_nodes" in redis_kwargs and "service_name" in redis_kwargs:
return _init_redis_sentinel(redis_kwargs)
@ -392,21 +401,6 @@ def get_redis_async_client(
**env_overrides,
) -> Union[async_redis.Redis, async_redis.RedisCluster]:
redis_kwargs = _get_redis_client_logic(**env_overrides)
if "url" in redis_kwargs and redis_kwargs["url"] is not None:
if connection_pool is not None:
return async_redis.Redis(connection_pool=connection_pool)
args = _get_redis_url_kwargs(client=async_redis.Redis.from_url)
url_kwargs = {}
for arg in redis_kwargs:
if arg in args:
url_kwargs[arg] = redis_kwargs[arg]
else:
verbose_logger.debug(
"REDIS: ignoring argument: {}. Not an allowed async_redis.Redis.from_url arg.".format(
arg
)
)
return async_redis.Redis.from_url(**url_kwargs)
if "startup_nodes" in redis_kwargs:
from redis.cluster import ClusterNode
@ -469,6 +463,22 @@ def get_redis_async_client(
return cluster_client
if "url" in redis_kwargs and redis_kwargs["url"] is not None:
if connection_pool is not None:
return async_redis.Redis(connection_pool=connection_pool)
args = _get_redis_url_kwargs(client=async_redis.Redis.from_url)
url_kwargs = {}
for arg in redis_kwargs:
if arg in args:
url_kwargs[arg] = redis_kwargs[arg]
else:
verbose_logger.debug(
"REDIS: ignoring argument: {}. Not an allowed async_redis.Redis.from_url arg.".format(
arg
)
)
return async_redis.Redis.from_url(**url_kwargs)
# Check for Redis Sentinel
if "sentinel_nodes" in redis_kwargs and "service_name" in redis_kwargs:
return _init_async_redis_sentinel(redis_kwargs)
@ -482,9 +492,15 @@ def get_redis_async_client(
)
def get_redis_connection_pool(**env_overrides):
def get_redis_connection_pool(
**env_overrides,
) -> Optional[async_redis.BlockingConnectionPool]:
redis_kwargs = _get_redis_client_logic(**env_overrides)
verbose_logger.debug("get_redis_connection_pool: redis_kwargs", redis_kwargs)
if "startup_nodes" in redis_kwargs:
return None
if "url" in redis_kwargs and redis_kwargs["url"] is not None:
pool_kwargs = {
"timeout": REDIS_CONNECTION_POOL_TIMEOUT,
@ -504,7 +520,6 @@ def get_redis_connection_pool(**env_overrides):
connection_class = async_redis.SSLConnection
redis_kwargs.pop("ssl", None)
redis_kwargs["connection_class"] = connection_class
redis_kwargs.pop("startup_nodes", None)
return async_redis.BlockingConnectionPool(
timeout=REDIS_CONNECTION_POOL_TIMEOUT, **redis_kwargs
)

View file

@ -240,10 +240,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
if key in ("max_tokens", "max_completion_tokens"):
responses_api_request["max_output_tokens"] = value
elif key == "tools" and value is not None:
responses_api_request[
"tools"
] = self._convert_tools_to_responses_format(
cast(List[Dict[str, Any]], value)
responses_api_request["tools"] = (
self._convert_tools_to_responses_format(
cast(List[Dict[str, Any]], value)
)
)
elif key == "response_format":
text_format = self._transform_response_format_to_text_format(value)
@ -696,6 +696,20 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
verbose_logger.debug(
f"Chat provider: image -> {converted}"
)
elif item_type == "file":
# Map Chat Completion file to Responses API input_file
# {"type": "file", "file": {"file_data": "...", "filename": "..."}}
# -> {"type": "input_file", "file_data": "...", "filename": "..."}
file_data = item.get("file", {})
converted = {"type": "input_file"}
if isinstance(file_data, dict):
for key in ["file_id", "file_data", "filename"]:
if key in file_data:
converted[key] = file_data[key]
result.append(converted)
verbose_logger.debug(
f"Chat provider: file -> {converted}"
)
elif item_type in [
"input_text",
"input_image",
@ -1058,9 +1072,9 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
)
if provider_specific_fields:
function_chunk[
"provider_specific_fields"
] = provider_specific_fields
function_chunk["provider_specific_fields"] = (
provider_specific_fields
)
tool_call_index = parsed_chunk.get("output_index", 0)
tool_call_chunk = ChatCompletionToolCallChunk(
@ -1133,9 +1147,9 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
# Add provider_specific_fields to function if present
if provider_specific_fields:
function_chunk[
"provider_specific_fields"
] = provider_specific_fields
function_chunk["provider_specific_fields"] = (
provider_specific_fields
)
tool_call_index = parsed_chunk.get("output_index", 0)
tool_call_chunk = ChatCompletionToolCallChunk(

View file

@ -83,7 +83,26 @@ class LangsmithLogger(CustomBatchLogger):
if _batch_size:
self.batch_size = int(_batch_size)
self.log_queue: List[LangsmithQueueObject] = []
asyncio.create_task(self.periodic_flush())
self._flush_task: Optional[asyncio.Task[Any]] = self._start_periodic_flush_task()
def _start_periodic_flush_task(self) -> Optional[asyncio.Task[Any]]:
"""Start the periodic flush task only when an event loop is already running."""
try:
loop = asyncio.get_running_loop()
except RuntimeError:
verbose_logger.debug(
"Langsmith logger init: no running event loop, skipping periodic flush task startup"
)
return None
return loop.create_task(self.periodic_flush())
def _ensure_periodic_flush_task(self) -> None:
# This helper is intentionally synchronous. In asyncio's cooperative
# execution model, there is no await between the check and assignment,
# so one caller cannot interleave here and create a duplicate task.
if self._flush_task is None or self._flush_task.done():
self._flush_task = self._start_periodic_flush_task()
def get_credentials_from_env(
self,
@ -266,6 +285,7 @@ class LangsmithLogger(CustomBatchLogger):
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
try:
self._ensure_periodic_flush_task()
sampling_rate = self._get_sampling_rate_to_use_for_request(kwargs=kwargs)
random_sample = random.random()
if random_sample > sampling_rate:
@ -307,17 +327,18 @@ class LangsmithLogger(CustomBatchLogger):
)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
sampling_rate = self._get_sampling_rate_to_use_for_request(kwargs=kwargs)
random_sample = random.random()
if random_sample > sampling_rate:
verbose_logger.info(
"Skipping Langsmith logging. Sampling rate={}, random_sample={}".format(
sampling_rate, random_sample
)
)
return # Skip logging
verbose_logger.info("Langsmith Failure Event Logging!")
try:
self._ensure_periodic_flush_task()
sampling_rate = self._get_sampling_rate_to_use_for_request(kwargs=kwargs)
random_sample = random.random()
if random_sample > sampling_rate:
verbose_logger.info(
"Skipping Langsmith logging. Sampling rate={}, random_sample={}".format(
sampling_rate, random_sample
)
)
return # Skip logging
verbose_logger.info("Langsmith Failure Event Logging!")
credentials = self._get_credentials_to_use_for_request(kwargs=kwargs)
data = self._prepare_log_data(
kwargs=kwargs,

View file

@ -64,6 +64,7 @@ _FINISH_REASON_MAP: dict[str, OpenAIChatCompletionFinishReason] = {
"end_turn": "stop",
"max_tokens": "length",
"tool_use": "tool_calls",
"refusal": "content_filter",
"compaction": "length",
# Cohere
"COMPLETE": "stop",

View file

@ -1498,17 +1498,49 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915
from litellm.types.llms.vertex_ai import BlobType
content_str: str = ""
inline_data: Optional[BlobType] = None
inline_data_list: List[BlobType] = []
if "content" in message:
if isinstance(message["content"], str):
content_str = message["content"]
# Detect data-URL images (e.g. from Anthropic tool_result with a single image block
# that was serialised as a plain string by translate_anthropic_messages_to_openai)
# and promote them to inline_data so Gemini receives actual image bytes.
if content_str[:5].lower() == "data:" and ";base64," in content_str:
try:
mime_rest = content_str[5:].split(";base64,", 1)
if len(mime_rest) == 2 and mime_rest[0].startswith("image/"):
# Strip any extra parameters (e.g. ";charset=UTF-8") from the MIME segment
clean_mime = mime_rest[0].split(";")[0].strip()
inline_data_list.append(
BlobType(data=mime_rest[1], mime_type=clean_mime)
)
content_str = ""
except Exception as e:
verbose_logger.warning(
f"Failed to parse data URL in tool response: {e}"
)
elif isinstance(message["content"], List):
content_list = message["content"]
for content in content_list:
content_type = content.get("type", "")
if content_type == "text":
content_str += content.get("text", "")
elif content_type == "image":
# Anthropic-native image block: {"type": "image", "source": {"type": "base64", ...}}
source = content.get("source", {})
if isinstance(source, dict) and source.get("type") == "base64":
try:
inline_data_list.append(
BlobType(
data=source.get("data", ""),
mime_type=source.get("media_type", "image/jpeg"),
)
)
except Exception as e:
verbose_logger.warning(
f"Failed to process Anthropic image block in tool response: {e}"
)
elif content_type in ("input_image", "image_url"):
# Extract image for inline_data (for Computer Use screenshots and tool results)
image_url_data = content.get("image_url", "")
@ -1524,9 +1556,11 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915
image_obj = convert_to_anthropic_image_obj(
image_url, format=None
)
inline_data = BlobType(
data=image_obj["data"],
mime_type=image_obj["media_type"],
inline_data_list.append(
BlobType(
data=image_obj["data"],
mime_type=image_obj["media_type"],
)
)
except Exception as e:
verbose_logger.warning(
@ -1551,9 +1585,11 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915
file_obj = convert_to_anthropic_image_obj(
file_data, format=None
)
inline_data = BlobType(
data=file_obj["data"],
mime_type=file_obj["media_type"],
inline_data_list.append(
BlobType(
data=file_obj["data"],
mime_type=file_obj["media_type"],
)
)
except Exception as e:
verbose_logger.warning(
@ -1607,13 +1643,12 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915
# Create part with function_response, and optionally inline_data for images (Computer Use)
_part: VertexPartType = {"function_response": _function_response}
# For Computer Use, if we have an image, we need separate parts:
# For Computer Use, if we have images/files, we need separate parts:
# - One part with function_response
# - One part with inline_data
# - One part per inline_data item
# Gemini's PartType is a oneof, so we can't have both in the same part
if inline_data:
image_part: VertexPartType = {"inline_data": inline_data}
return [_part, image_part]
if inline_data_list:
return [_part] + [{"inline_data": d} for d in inline_data_list]
return _part

View file

@ -131,14 +131,9 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config):
if result_effort == "none" and not supports_none:
result.pop("reasoning_effort")
# Azure Chat Completions: gpt-5.4+ does not support tools + reasoning together.
# Drop reasoning_effort when both are present (OpenAI routes to Responses API; Azure does not).
if self.is_model_gpt_5_4_plus_model(model):
has_tools = bool(
non_default_params.get("tools") or optional_params.get("tools")
)
if has_tools and result_effort not in (None, "none"):
result.pop("reasoning_effort", None)
# Azure gpt-5.4+ with tools + reasoning_effort is now routed to the
# Responses API bridge (same as OpenAI), so we no longer need to drop
# reasoning_effort here. See: responses_api_bridge_check() in main.py.
return result

View file

@ -185,11 +185,16 @@ class FireworksAIConfig(OpenAIGPTConfig):
): # allow user to toggle this feature.
return content
if isinstance(content["image_url"], str):
content["image_url"] = f"{content['image_url']}#transform=inline"
# Skip base64 data URLs — appending #transform=inline corrupts the
# base64 payload and causes an "Incorrect padding" decode error on
# the Fireworks side. Data URLs are already inlined by definition.
# Lower-case before checking: URI schemes are case-insensitive (RFC 3986).
if not content["image_url"].lower().startswith("data:"):
content["image_url"] = f"{content['image_url']}#transform=inline"
elif isinstance(content["image_url"], dict):
content["image_url"][
"url"
] = f"{content['image_url']['url']}#transform=inline"
url = content["image_url"]["url"]
if not url.lower().startswith("data:"):
content["image_url"]["url"] = f"{url}#transform=inline"
return content
def _transform_tools(

View file

@ -148,5 +148,12 @@ class MistralAudioTranscriptionConfig(BaseAudioTranscriptionConfig):
text = response_json.get("text") or ""
response = TranscriptionResponse(text=text)
# Preserve Mistral-specific fields (e.g. diarization segments)
if "segments" in response_json:
response["segments"] = response_json["segments"]
if "language" in response_json:
response["language"] = response_json["language"]
response._hidden_params = response_json
return response

View file

@ -51,6 +51,7 @@ class ContextCachingEndpoints(VertexBase):
vertex_project: Optional[str],
vertex_location: Optional[str],
vertex_auth_header: Optional[str],
model: Optional[str] = None,
) -> Tuple[Optional[str], str]:
"""
Internal function. Returns the token and url for the call.
@ -89,7 +90,7 @@ class ContextCachingEndpoints(VertexBase):
stream=None,
auth_header=auth_header,
url=url,
model=None,
model=model,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_api_version="v1beta1"
@ -109,6 +110,7 @@ class ContextCachingEndpoints(VertexBase):
vertex_project: Optional[str],
vertex_location: Optional[str],
vertex_auth_header: Optional[str],
model: Optional[str] = None,
) -> Optional[str]:
"""
Checks if content already cached.
@ -128,6 +130,7 @@ class ContextCachingEndpoints(VertexBase):
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header,
model=model,
)
page_token: Optional[str] = None
@ -201,6 +204,7 @@ class ContextCachingEndpoints(VertexBase):
vertex_project: Optional[str],
vertex_location: Optional[str],
vertex_auth_header: Optional[str],
model: Optional[str] = None,
) -> Optional[str]:
"""
Checks if content already cached.
@ -220,6 +224,7 @@ class ContextCachingEndpoints(VertexBase):
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header,
model=model,
)
page_token: Optional[str] = None
@ -342,6 +347,7 @@ class ContextCachingEndpoints(VertexBase):
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header,
model=model,
)
headers = {
@ -377,6 +383,7 @@ class ContextCachingEndpoints(VertexBase):
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header,
model=model,
)
if google_cache_name:
return non_cached_messages, optional_params, google_cache_name
@ -488,6 +495,7 @@ class ContextCachingEndpoints(VertexBase):
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header,
model=model,
)
headers = {
@ -520,6 +528,7 @@ class ContextCachingEndpoints(VertexBase):
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_auth_header=vertex_auth_header,
model=model,
)
if google_cache_name:

View file

@ -3079,6 +3079,16 @@ class ModelResponseIterator:
)
model_response.choices.append(choice)
# Also handle the case where the final chunk has empty
# content (e.g. text:"") WITH finishReason. In this case
# _process_candidates DOES create a choice, but maps
# finishReason="STOP" to "stop" because the current chunk
# has no tool_calls. Override if we saw tool_calls earlier.
if self.has_seen_tool_calls:
for choice in model_response.choices:
if choice.finish_reason == "stop":
choice.finish_reason = "tool_calls"
setattr(model_response, "vertex_ai_grounding_metadata", grounding_metadata) # type: ignore
setattr(model_response, "vertex_ai_url_context_metadata", url_context_metadata) # type: ignore
setattr(model_response, "vertex_ai_safety_ratings", safety_ratings) # type: ignore

View file

@ -105,12 +105,25 @@ class VertexAIPartnerModelsTokenCounter(VertexBase):
# Extract Vertex AI credentials and settings
vertex_credentials = self.get_vertex_ai_credentials(litellm_params)
vertex_project = self.get_vertex_ai_project(litellm_params)
vertex_location = self.get_vertex_ai_location(litellm_params)
# Map empty location/cluade models to a supported region for count-tokens endpoint
# Check for count_tokens specific location override
vertex_count_tokens_location = litellm_params.get("vertex_count_tokens_location")
vertex_location_raw = self.get_vertex_ai_location(litellm_params)
# Determine final location with precedence:
# 1. vertex_count_tokens_location (if provided)
# 2. vertex_location (if provided)
# 3. Default to us-east5 for Claude models when no location is set
# Supported regions: us-east5, europe-west1, asia-southeast1
# https://docs.cloud.google.com/vertex-ai/generative-ai/docs/partner-models/claude/count-tokens
if not vertex_location or "claude" in model.lower():
vertex_location = "us-central1"
if vertex_count_tokens_location:
vertex_location: str = vertex_count_tokens_location
elif vertex_location_raw:
vertex_location = vertex_location_raw
elif "claude" in model.lower():
vertex_location = "us-east5"
else:
vertex_location = "us-east5"
# Get access token and resolved project ID
access_token, project_id = await self._ensure_access_token_async(

View file

@ -955,16 +955,6 @@ def responses_api_bridge_check(
model_info["mode"] = "responses"
model = model.replace("responses/", "")
# OpenAI gpt-5.4+ chat-completions calls with both tools + reasoning_effort
# must be bridged to Responses API.
if (
custom_llm_provider == "openai"
and OpenAIGPT5Config.is_model_gpt_5_4_plus_model(model)
and tools
and reasoning_effort is not None
):
model_info["mode"] = "responses"
model = model.replace("responses/", "")
except Exception as e:
verbose_logger.debug("Error getting model info: {}".format(e))
@ -974,6 +964,19 @@ def responses_api_bridge_check(
model = model.replace("responses/", "")
mode = "responses"
model_info["mode"] = mode
# OpenAI/Azure gpt-5.4+ chat-completions calls with both tools + reasoning_effort
# must be bridged to Responses API.
if (
custom_llm_provider in ("openai", "azure")
and OpenAIGPT5Config.is_model_gpt_5_4_plus_model(model)
and tools
and reasoning_effort is not None
and model_info.get("mode") != "responses"
):
model_info["mode"] = "responses"
model = model.replace("responses/", "")
return model_info, model

View file

@ -6152,7 +6152,8 @@
"max_query_tokens": 4096,
"max_tokens": 32768,
"mode": "rerank",
"output_cost_per_token": 0.0
"output_cost_per_token": 0.0,
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-cohere-rerank-4-0-in-microsoft-foundry/4477076"
},
"azure_ai/cohere-rerank-v4.0-fast": {
"input_cost_per_query": 0.002,
@ -6163,7 +6164,8 @@
"max_query_tokens": 4096,
"max_tokens": 32768,
"mode": "rerank",
"output_cost_per_token": 0.0
"output_cost_per_token": 0.0,
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-cohere-rerank-4-0-in-microsoft-foundry/4477076"
},
"azure_ai/deepseek-v3.2": {
"input_cost_per_token": 5.8e-07,
@ -6173,6 +6175,7 @@
"max_tokens": 163840,
"mode": "chat",
"output_cost_per_token": 1.68e-06,
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-deepseek-v3-2-and-deepseek-v3-2-speciale-in-microsoft-foundry/4477549",
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
@ -6187,6 +6190,7 @@
"max_tokens": 163840,
"mode": "chat",
"output_cost_per_token": 1.68e-06,
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-deepseek-v3-2-and-deepseek-v3-2-speciale-in-microsoft-foundry/4477549",
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
@ -16936,6 +16940,18 @@
"supports_system_messages": true,
"supports_tool_choice": true
},
"gpt-4-0314": {
"deprecation_date": "2026-03-26",
"input_cost_per_token": 3e-05,
"litellm_provider": "openai",
"max_input_tokens": 8192,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_token": 6e-05,
"supports_system_messages": true,
"supports_tool_choice": true
},
"gpt-4-0613": {
"deprecation_date": "2025-06-06",
"input_cost_per_token": 3e-05,
@ -30506,7 +30522,7 @@
"output_cost_per_token": 5.4e-06,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models",
"supported_regions": [
"us-west2"
"us-central1"
],
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -30526,7 +30542,7 @@
"output_cost_per_token_batches": 8.4e-07,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models",
"supported_regions": [
"us-west2"
"global"
],
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -30543,6 +30559,9 @@
"mode": "chat",
"output_cost_per_token": 5.4e-06,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models",
"supported_regions": [
"us-central1"
],
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
@ -31167,7 +31186,10 @@
"input_cost_per_token": 3e-07,
"output_cost_per_token": 1.2e-06,
"ocr_cost_per_page": 0.0003,
"source": "https://cloud.google.com/vertex-ai/pricing"
"source": "https://cloud.google.com/vertex-ai/pricing",
"supported_regions": [
"us-central1"
]
},
"vertex_ai/openai/gpt-oss-120b-maas": {
"input_cost_per_token": 1.5e-07,

View file

@ -0,0 +1,37 @@
from typing import TYPE_CHECKING
from litellm.types.guardrails import SupportedGuardrailIntegrations
from .akto import AktoGuardrail
if TYPE_CHECKING:
from litellm.types.guardrails import Guardrail, LitellmParams
def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"):
import litellm
_akto_callback = AktoGuardrail(
akto_base_url=getattr(litellm_params, "akto_base_url", None),
akto_api_key=getattr(litellm_params, "akto_api_key", None),
akto_account_id=getattr(litellm_params, "akto_account_id", None),
akto_vxlan_id=getattr(litellm_params, "akto_vxlan_id", None),
unreachable_fallback=getattr(litellm_params, "unreachable_fallback", "fail_closed"),
guardrail_timeout=getattr(litellm_params, "guardrail_timeout", None),
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
)
litellm.logging_callback_manager.add_litellm_callback(_akto_callback)
return _akto_callback
guardrail_initializer_registry = {
SupportedGuardrailIntegrations.AKTO.value: initialize_guardrail,
}
guardrail_class_registry = {
SupportedGuardrailIntegrations.AKTO.value: AktoGuardrail,
}

View file

@ -0,0 +1,456 @@
"""Akto guardrail integration for LiteLLM proxy.
Uses a two-config-entry pattern:
- akto-validate (pre_call): Checks request against Akto guardrails, blocks if flagged.
- akto-ingest (post_call): Sends request+response to Akto for data ingestion.
For monitor-only mode, enable only akto-ingest without akto-validate.
"""
import asyncio
import json
import os
from datetime import datetime
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple, Type
from fastapi import HTTPException
import httpx
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
HTTP_PROXY_PATH = "/api/http-proxy"
AKTO_CONNECTOR_NAME = "litellm"
DEFAULT_GUARDRAIL_TIMEOUT = 5
class AktoGuardrail(CustomGuardrail):
"""LiteLLM guardrail hook that validates and ingests LLM traffic via the Akto API."""
# Maps event_hook to the input_type it should handle; mismatches are no-ops
HOOK_TO_INPUT = {"pre_call": "request", "post_call": "response"}
@staticmethod
def get_config_model() -> Type["GuardrailConfigModel"]:
"""Return the Pydantic config model for YAML-based initialization."""
from litellm.types.proxy.guardrails.guardrail_hooks.akto import (
AktoConfigModel,
)
return AktoConfigModel
def __init__(
self,
akto_base_url: Optional[str] = None,
akto_api_key: Optional[str] = None,
akto_account_id: Optional[str] = None,
akto_vxlan_id: Optional[str] = None,
unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
guardrail_timeout: Optional[int] = None,
**kwargs: Any,
) -> None:
"""Initialize the Akto guardrail.
Args:
akto_base_url: Akto API base URL. Falls back to AKTO_GUARDRAIL_API_BASE env var.
akto_api_key: Akto API key. Falls back to AKTO_API_KEY env var.
akto_account_id: Akto account ID. Falls back to AKTO_ACCOUNT_ID env var, then "1000000".
akto_vxlan_id: Akto VXLAN ID. Falls back to AKTO_VXLAN_ID env var, then "0".
unreachable_fallback: Behavior when Akto is unreachable — block or allow.
guardrail_timeout: HTTP timeout in seconds for Akto API calls.
"""
self.async_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback,
)
self.background_tasks: set = set()
self.akto_base_url = (akto_base_url or os.environ.get("AKTO_GUARDRAIL_API_BASE", "")).rstrip("/")
if not self.akto_base_url:
raise ValueError("akto_base_url is required. Set AKTO_GUARDRAIL_API_BASE or pass it in litellm_params.")
self.akto_api_key = akto_api_key or os.environ.get("AKTO_API_KEY", "")
if not self.akto_api_key:
raise ValueError("akto_api_key is required. Set AKTO_API_KEY or pass it in litellm_params.")
self.unreachable_fallback: Literal["fail_closed", "fail_open"] = unreachable_fallback
self.guardrail_timeout = guardrail_timeout or DEFAULT_GUARDRAIL_TIMEOUT
self.akto_account_id = akto_account_id or os.environ.get("AKTO_ACCOUNT_ID", "1000000")
self.akto_vxlan_id = akto_vxlan_id or os.environ.get("AKTO_VXLAN_ID", "0")
kwargs["supported_event_hooks"] = [
GuardrailEventHooks.pre_call,
GuardrailEventHooks.post_call,
]
super().__init__(**kwargs)
verbose_proxy_logger.debug(
"Akto guardrail initialized: base_url=%s fallback=%s",
self.akto_base_url,
self.unreachable_fallback,
)
@staticmethod
def resolve_metadata_value(request_data: Optional[dict], key: str) -> Optional[str]:
"""Look up a metadata value from litellm_metadata or metadata dicts."""
if request_data is None:
return None
for dict_key in ("litellm_metadata", "metadata"):
container = request_data.get(dict_key) or {}
if isinstance(container, dict) and container:
value = container.get(key)
if value is not None:
return str(value).strip()
return None
@staticmethod
def extract_request_path(request_data: dict) -> str:
"""Extract the API route from request metadata, defaulting to /v1/chat/completions."""
metadata = request_data.get("metadata") or {}
if not isinstance(metadata, dict):
metadata = {}
route = metadata.get("user_api_key_request_route")
return route if route else "/v1/chat/completions"
def prepare_headers(self) -> Dict[str, str]:
"""Build HTTP headers for the Akto API call."""
return {
"content-type": "application/json",
"Authorization": self.akto_api_key,
}
@staticmethod
def build_query_params(*, guardrails: bool, ingest_data: bool) -> Dict[str, str]:
"""Build query params that control Akto backend behavior (guardrail check and/or data ingestion)."""
params: Dict[str, str] = {"akto_connector": AKTO_CONNECTOR_NAME}
if guardrails:
params["guardrails"] = "true"
if ingest_data:
params["ingest_data"] = "true"
return params
@staticmethod
def build_request_headers(request_data: dict) -> Dict[str, str]:
"""Build the requestHeaders field from proxy request headers."""
headers: Dict[str, str] = {"content-type": "application/json"}
proxy_req = request_data.get("proxy_server_request", {})
if not isinstance(proxy_req, dict):
return headers
proxy_req_headers = proxy_req.get("headers")
if isinstance(proxy_req_headers, dict):
for key, val in proxy_req_headers.items():
if key and val:
headers[str(key).lower()] = str(val)
return headers
@staticmethod
def build_request_body(
inputs: GenericGuardrailAPIInputs,
request_data: Optional[dict] = None,
) -> Dict[str, Any]:
"""Build the LLM request body from guardrail inputs (messages, model, tools)."""
model = inputs.get("model", "") or ""
body: Dict[str, Any] = {"model": model}
structured = inputs.get("structured_messages")
if structured:
body["messages"] = structured
elif request_data is not None and request_data.get("messages"):
body["messages"] = request_data["messages"]
if request_data.get("model"):
body["model"] = request_data["model"]
else:
texts = inputs.get("texts", [])
body["messages"] = [{"role": "user", "content": t} for t in texts] if texts else []
tools = inputs.get("tools")
if tools:
body["tools"] = tools
elif request_data is not None and request_data.get("tools"):
body["tools"] = request_data["tools"]
tool_calls = inputs.get("tool_calls")
if tool_calls:
body["tool_calls"] = tool_calls
return body
@staticmethod
def build_response_body(
inputs: GenericGuardrailAPIInputs,
request_data: Optional[dict] = None,
) -> Dict[str, Any]:
"""Build the LLM response body, preferring the actual model response if available."""
model_response = request_data.get("response") if request_data else None
if model_response is not None and hasattr(model_response, "model_dump"):
return model_response.model_dump()
texts = inputs.get("texts", [])
if texts:
return {"choices": [{"message": {"content": t, "role": "assistant"}} for t in texts]}
return {}
@staticmethod
def build_tag_metadata(request_data: dict) -> Dict[str, str]:
"""Build tag/metadata dict with user_id and team_id for Akto tracking."""
tag: Dict[str, str] = {"gen-ai": "Gen AI"}
user_id = AktoGuardrail.resolve_metadata_value(request_data, "user_api_key_user_id")
team_id = AktoGuardrail.resolve_metadata_value(request_data, "user_api_key_team_id")
if user_id:
tag["user_id"] = user_id
if team_id:
tag["team_id"] = team_id
return tag
def build_akto_payload(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
*,
status_code: int = 200,
include_response: bool = False,
) -> Dict[str, Any]:
"""Build the flat MIRRORING payload sent to Akto's HTTP proxy endpoint.
All body fields use double-encoding: json.dumps({"body": json.dumps(actual_body)})
to match the canonical CLI hook format.
"""
request_path = self.extract_request_path(request_data)
request_headers = self.build_request_headers(request_data)
request_body = self.build_request_body(inputs, request_data)
tag = self.build_tag_metadata(request_data)
response_payload = json.dumps({}) # Empty body wrapper when no response yet
response_headers: Dict[str, str] = {}
if include_response:
response_body = self.build_response_body(inputs, request_data)
response_payload = json.dumps({"body": json.dumps(response_body)}) # Double-encoded
response_headers = {"content-type": "application/json"}
# Extract client IP from proxy headers
ip = ""
proxy_req = request_data.get("proxy_server_request", {})
proxy_headers = proxy_req.get("headers", {}) if isinstance(proxy_req, dict) else {}
if isinstance(proxy_headers, dict):
ip = proxy_headers.get("x-forwarded-for") or proxy_headers.get("x-real-ip") or ""
if "," in ip:
ip = ip.split(",")[0].strip()
return {
"path": request_path,
"requestHeaders": json.dumps(request_headers),
"responseHeaders": json.dumps(response_headers),
"method": "POST",
"requestPayload": json.dumps({"body": json.dumps(request_body)}), # Double-encoded
"responsePayload": response_payload,
"ip": ip,
"destIp": "127.0.0.1",
"time": str(int(datetime.now().timestamp() * 1000)),
"statusCode": str(status_code),
"type": "HTTP/1.1",
"status": str(status_code),
"akto_account_id": self.akto_account_id,
"akto_vxlan_id": self.akto_vxlan_id,
"is_pending": "false",
"source": "MIRRORING",
"direction": None,
"process_id": None,
"socket_id": None,
"daemonset_id": None,
"enabled_graph": None,
"tag": json.dumps(tag),
"metadata": json.dumps(tag),
"contextSource": "AGENTIC",
}
async def send_request(
self,
*,
guardrails: bool,
ingest_data: bool,
payload: dict,
) -> httpx.Response:
"""Send an HTTP POST to the Akto API endpoint."""
endpoint = f"{self.akto_base_url}{HTTP_PROXY_PATH}"
params = self.build_query_params(guardrails=guardrails, ingest_data=ingest_data)
headers = self.prepare_headers()
return await self.async_handler.post(
url=endpoint,
data=json.dumps(payload),
params=params,
headers=headers,
timeout=self.guardrail_timeout,
)
@staticmethod
def handle_guardrail_response(response: httpx.Response) -> Tuple[bool, str]:
"""Parse the Akto guardrail response. Returns (allowed, reason)."""
if response.status_code != 200:
verbose_proxy_logger.error("Akto returned HTTP %d", response.status_code)
raise httpx.HTTPStatusError(
f"Akto returned unexpected status {response.status_code}",
request=response.request,
response=response,
)
try:
result = response.json()
except (json.JSONDecodeError, ValueError) as e:
response_text = getattr(response, "text", "")
verbose_proxy_logger.error(
"Akto returned non-JSON body for status 200: %r",
response_text[:200],
)
raise httpx.RequestError(
"Akto returned non-JSON body",
request=response.request,
) from e
if not isinstance(result, dict):
return True, ""
data = result.get("data") or {}
if not isinstance(data, dict):
return True, ""
guardrails_result = data.get("guardrailsResult") or {}
if not isinstance(guardrails_result, dict):
return True, ""
return (
bool(guardrails_result.get("Allowed", True)),
str(guardrails_result.get("Reason", "")),
)
def handle_unreachable(
self,
inputs: GenericGuardrailAPIInputs,
error: Exception,
) -> GenericGuardrailAPIInputs:
"""Handle Akto being unreachable based on fail_open/fail_closed config."""
if self.unreachable_fallback == "fail_open":
verbose_proxy_logger.critical(
"Akto unreachable (fail-open): %s",
str(error),
exc_info=error,
)
return inputs
verbose_proxy_logger.error("Akto unreachable (fail-closed): %s", str(error))
raise HTTPException(
status_code=503,
detail="Akto guardrail service unreachable",
)
async def fire_and_forget_request(
self,
*,
guardrails: bool,
ingest_data: bool,
payload: dict,
) -> None:
"""Send a request without awaiting it in the caller. Errors are logged, not raised."""
try:
response = await self.send_request(
guardrails=guardrails,
ingest_data=ingest_data,
payload=payload,
)
if response.status_code != 200:
verbose_proxy_logger.error(
"Akto fire-and-forget returned HTTP %d",
response.status_code,
)
except Exception as e:
verbose_proxy_logger.error("Akto fire-and-forget error: %s", str(e))
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj=None,
) -> GenericGuardrailAPIInputs:
"""Main entry point called by LiteLLM's guardrail framework.
Pre_call (input_type="request"):
- Awaits guardrail check. If blocked, fires off ingest with 403 marker and raises.
Post_call (input_type="response"):
- Fire-and-forget combined guardrail + ingest call.
"""
# Skip if this hook doesn't handle the current input_type
expected = self.HOOK_TO_INPUT.get(str(self.event_hook))
if expected and expected != input_type:
return inputs
if input_type == "request":
# Pre_call: awaited guardrail check (no ingestion)
payload = self.build_akto_payload(inputs, request_data, include_response=False)
try:
response = await self.send_request(
guardrails=True,
ingest_data=False,
payload=payload,
)
allowed, reason = self.handle_guardrail_response(response)
except HTTPException:
raise
except (httpx.RequestError, httpx.HTTPStatusError) as e:
return self.handle_unreachable(
inputs=inputs,
error=e,
)
if not allowed:
# Build a blocked marker payload with 403 status and reason
blocked_payload = self.build_akto_payload(
inputs,
request_data,
include_response=False,
status_code=403,
)
blocked_payload["responsePayload"] = json.dumps(
{
"body": json.dumps({"x-blocked-by": "Akto Proxy", "reason": reason}),
}
)
blocked_payload["responseHeaders"] = json.dumps(
{"content-type": "application/json"},
)
# Fire-and-forget ingest of the blocked request, then raise 403
task = asyncio.create_task(
self.fire_and_forget_request(
guardrails=False,
ingest_data=True,
payload=blocked_payload,
)
)
self.background_tasks.add(task)
task.add_done_callback(self.background_tasks.discard)
raise HTTPException(
status_code=403,
detail=reason or "Blocked by Akto Guardrails",
)
elif input_type == "response":
# Post_call: fire-and-forget combined guardrail + ingest
payload = self.build_akto_payload(inputs, request_data, include_response=True)
task = asyncio.create_task(
self.fire_and_forget_request(
guardrails=True,
ingest_data=True,
payload=payload,
)
)
self.background_tasks.add(task)
task.add_done_callback(self.background_tasks.discard)
return inputs

View file

@ -639,9 +639,9 @@ except ImportError:
server_root_path = get_server_root_path()
_license_check = LicenseCheck()
premium_user: bool = _license_check.is_premium()
premium_user_data: Optional[
"EnterpriseLicenseData"
] = _license_check.airgapped_license_data
premium_user_data: Optional["EnterpriseLicenseData"] = (
_license_check.airgapped_license_data
)
global_max_parallel_request_retries_env: Optional[str] = os.getenv(
"LITELLM_GLOBAL_MAX_PARALLEL_REQUEST_RETRIES"
)
@ -1524,9 +1524,9 @@ master_key: Optional[str] = None
config_agents: Optional[List[AgentConfig]] = None
otel_logging = False
prisma_client: Optional[PrismaClient] = None
shared_aiohttp_session: Optional[
"ClientSession"
] = None # Global shared session for connection reuse
shared_aiohttp_session: Optional["ClientSession"] = (
None # Global shared session for connection reuse
)
user_api_key_cache = DualCache(
default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value
)
@ -1534,13 +1534,13 @@ model_max_budget_limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(
dual_cache=user_api_key_cache
)
litellm.logging_callback_manager.add_litellm_callback(model_max_budget_limiter)
redis_usage_cache: Optional[
RedisCache
] = None # redis cache used for tracking spend, tpm/rpm limits
redis_usage_cache: Optional[RedisCache] = (
None # redis cache used for tracking spend, tpm/rpm limits
)
polling_via_cache_enabled: Union[Literal["all"], List[str], bool] = False
native_background_mode: List[
str
] = [] # Models that should use native provider background mode instead of polling
native_background_mode: List[str] = (
[]
) # Models that should use native provider background mode instead of polling
polling_cache_ttl: int = 3600 # Default 1 hour TTL for polling cache
user_custom_auth = None
user_custom_key_generate = None
@ -1900,9 +1900,9 @@ async def update_cache( # noqa: PLR0915
_id = "team_id:{}".format(team_id)
try:
# Fetch the existing cost for the given user
existing_spend_obj: Optional[
LiteLLM_TeamTable
] = await user_api_key_cache.async_get_cache(key=_id)
existing_spend_obj: Optional[LiteLLM_TeamTable] = (
await user_api_key_cache.async_get_cache(key=_id)
)
if existing_spend_obj is None:
# do nothing if team not in api key cache
return
@ -2023,11 +2023,9 @@ def run_ollama_serve():
with open(os.devnull, "w") as devnull:
subprocess.Popen(command, stdout=devnull, stderr=devnull)
except Exception as e:
verbose_proxy_logger.debug(
f"""
verbose_proxy_logger.debug(f"""
LiteLLM Warning: proxy started with `ollama` model\n`ollama serve` failed with Exception{e}. \nEnsure you run `ollama serve`
"""
)
""")
def _get_process_rss_mb() -> Optional[float]:
@ -3321,7 +3319,7 @@ class ProxyConfig:
async_only_mode=True # only init async clients
),
ignore_invalid_deployments=True, # don't raise an error if a deployment is invalid
) # type:ignore
) # type: ignore
if redis_usage_cache is not None and router.cache.redis_cache is None:
router._update_redis_cache(cache=redis_usage_cache)
@ -4978,10 +4976,10 @@ class ProxyConfig:
)
try:
guardrails_in_db: List[
Guardrail
] = await GuardrailRegistry.get_all_guardrails_from_db(
prisma_client=prisma_client
guardrails_in_db: List[Guardrail] = (
await GuardrailRegistry.get_all_guardrails_from_db(
prisma_client=prisma_client
)
)
verbose_proxy_logger.debug(
"guardrails from the DB %s", str(guardrails_in_db)
@ -5363,9 +5361,9 @@ async def initialize( # noqa: PLR0915
user_api_base = api_base
dynamic_config[user_model]["api_base"] = api_base
if api_version:
os.environ[
"AZURE_API_VERSION"
] = api_version # set this for azure - litellm can read this from the env
os.environ["AZURE_API_VERSION"] = (
api_version # set this for azure - litellm can read this from the env
)
if max_tokens: # model-specific param
dynamic_config[user_model]["max_tokens"] = max_tokens
if temperature: # model-specific param
@ -5383,8 +5381,8 @@ async def initialize( # noqa: PLR0915
litellm.add_function_to_prompt = True
dynamic_config["general"]["add_function_to_prompt"] = True
if max_budget: # litellm-specific param
litellm.max_budget = max_budget
dynamic_config["general"]["max_budget"] = max_budget
litellm.max_budget = float(max_budget)
dynamic_config["general"]["max_budget"] = litellm.max_budget
if experimental:
pass
user_telemetry = telemetry
@ -5702,9 +5700,9 @@ class ProxyStartupEvent:
"""
from litellm.secret_managers.main import str_to_bool
_use_redis_transaction_buffer: Optional[
Union[bool, str]
] = general_settings.get("use_redis_transaction_buffer", False)
_use_redis_transaction_buffer: Optional[Union[bool, str]] = (
general_settings.get("use_redis_transaction_buffer", False)
)
if isinstance(_use_redis_transaction_buffer, str):
_use_redis_transaction_buffer = str_to_bool(_use_redis_transaction_buffer)
@ -12299,9 +12297,9 @@ async def get_config_list(
hasattr(sub_field_info, "description")
and sub_field_info.description is not None
):
nested_fields[
idx
].field_description = sub_field_info.description
nested_fields[idx].field_description = (
sub_field_info.description
)
idx += 1
_stored_in_db = None

View file

@ -17,6 +17,9 @@ from litellm.types.proxy.guardrails.guardrail_hooks.grayswan import (
from litellm.types.proxy.guardrails.guardrail_hooks.ibm import (
IBMGuardrailsBaseConfigModel,
)
from litellm.types.proxy.guardrails.guardrail_hooks.akto import (
AktoConfigModel,
)
from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import (
ContentFilterCategoryConfig,
)
@ -33,7 +36,7 @@ Pydantic object defining how to set guardrails on litellm proxy
guardrails:
- guardrail_name: "bedrock-pre-guard"
litellm_params:
guardrail: bedrock # supported values: "aporia", "bedrock", "lakera", "zscaler_ai_guard"
guardrail: bedrock # supported values: "akto", "aporia", "bedrock", "lakera", "zscaler_ai_guard"
mode: "during_call"
guardrailIdentifier: ff6ujrregl1q
guardrailVersion: "DRAFT"
@ -79,6 +82,7 @@ class SupportedGuardrailIntegrations(Enum):
SEMANTIC_GUARD = "semantic_guard"
MCP_END_USER_PERMISSION = "mcp_end_user_permission"
BLOCK_CODE_EXECUTION = "block_code_execution"
AKTO = "akto"
MCP_JWT_SIGNER = "mcp_jwt_signer"
@ -737,6 +741,7 @@ class LitellmParams(
NomaGuardrailConfigModel,
ToolPermissionGuardrailConfigModel,
ZscalerAIGuardConfigModel,
AktoConfigModel,
JavelinGuardrailConfigModel,
BaseLitellmParams,
EnkryptAIGuardrailConfigs,

View file

@ -0,0 +1,55 @@
from typing import Optional, Literal
from pydantic import Field
from .base import GuardrailConfigModel
class AktoConfigModel(GuardrailConfigModel):
"""
Config for the Akto guardrail.
Use two separate config entries to control behaviour:
akto-validate (mode: pre_call) -> check guardrails, block if flagged
akto-ingest (mode: post_call) -> ingest request+response data
"""
akto_base_url: Optional[str] = Field(
default=None,
description="Akto Guardrail API Base URL. Env: AKTO_GUARDRAIL_API_BASE.",
json_schema_extra={
"examples": [
"http://localhost:9090",
"https://akto-ingestion.example.com",
]
},
)
akto_api_key: Optional[str] = Field(
default=None,
description="API key for Akto. Env: AKTO_API_KEY.",
)
akto_account_id: Optional[str] = Field(
default=None,
description="Akto account ID for multi-tenant deployments. Env: AKTO_ACCOUNT_ID. Default: '1000000'.",
)
akto_vxlan_id: Optional[str] = Field(
default=None,
description="Akto VXLAN ID. Env: AKTO_VXLAN_ID. Default: '0'.",
)
unreachable_fallback: Literal["fail_closed", "fail_open"] = Field(
default="fail_closed",
description="What to do when Akto is unreachable. 'fail_open' = allow, 'fail_closed' = block.",
)
guardrail_timeout: Optional[int] = Field(
default=None,
description="HTTP timeout in seconds. Default: 5.",
)
@staticmethod
def ui_friendly_name() -> str:
return "Akto"

View file

@ -16940,6 +16940,18 @@
"supports_system_messages": true,
"supports_tool_choice": true
},
"gpt-4-0314": {
"deprecation_date": "2026-03-26",
"input_cost_per_token": 3e-05,
"litellm_provider": "openai",
"max_input_tokens": 8192,
"max_output_tokens": 4096,
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_token": 6e-05,
"supports_system_messages": true,
"supports_tool_choice": true
},
"gpt-4-0613": {
"deprecation_date": "2025-06-06",
"input_cost_per_token": 3e-05,
@ -30510,7 +30522,7 @@
"output_cost_per_token": 5.4e-06,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models",
"supported_regions": [
"us-west2"
"us-central1"
],
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -30530,7 +30542,7 @@
"output_cost_per_token_batches": 8.4e-07,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models",
"supported_regions": [
"us-west2"
"global"
],
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -30547,6 +30559,9 @@
"mode": "chat",
"output_cost_per_token": 5.4e-06,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models",
"supported_regions": [
"us-central1"
],
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
@ -31171,7 +31186,10 @@
"input_cost_per_token": 3e-07,
"output_cost_per_token": 1.2e-06,
"ocr_cost_per_page": 0.0003,
"source": "https://cloud.google.com/vertex-ai/pricing"
"source": "https://cloud.google.com/vertex-ai/pricing",
"supported_regions": [
"us-central1"
]
},
"vertex_ai/openai/gpt-oss-120b-maas": {
"input_cost_per_token": 1.5e-07,

View file

@ -0,0 +1,550 @@
import asyncio
import json
import os
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from starlette.exceptions import HTTPException
from litellm.types.utils import GenericGuardrailAPIInputs
from litellm.proxy.guardrails.guardrail_registry import guardrail_initializer_registry, guardrail_class_registry
from litellm.proxy.guardrails.guardrail_hooks.akto.akto import AktoGuardrail
# ---------------------------------------------------------------------------
# Registry tests
# ---------------------------------------------------------------------------
def test_akto_in_guardrail_initializer_registry():
assert "akto" in guardrail_initializer_registry
def test_akto_in_guardrail_class_registry():
assert "akto" in guardrail_class_registry
assert guardrail_class_registry["akto"] is AktoGuardrail
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture
def akto_validate():
"""AktoGuardrail configured for pre_call (akto-validate)."""
return AktoGuardrail(
akto_base_url="http://localhost:9090",
akto_api_key="test-token",
unreachable_fallback="fail_closed",
guardrail_name="test-akto-validate",
event_hook="pre_call",
)
@pytest.fixture
def akto_ingest():
"""AktoGuardrail configured for post_call (akto-ingest)."""
return AktoGuardrail(
akto_base_url="http://localhost:9090",
akto_api_key="test-token",
unreachable_fallback="fail_open",
guardrail_name="test-akto-ingest",
event_hook="post_call",
)
@pytest.fixture
def sample_inputs() -> GenericGuardrailAPIInputs:
return GenericGuardrailAPIInputs(
texts=["Hello, how are you?"],
model="gpt-4",
)
@pytest.fixture
def sample_request_data() -> dict:
return {
"metadata": {
"user_api_key_request_route": "/v1/chat/completions",
"user_api_key": "sk-test-123",
"user_api_key_user_id": "user-1",
"user_api_key_team_id": "team-1",
},
"proxy_server_request": {
"headers": {
"x-forwarded-for": "10.0.0.1",
}
},
}
def _mock_allowed_response():
mock = MagicMock(spec=httpx.Response)
mock.status_code = 200
mock.json.return_value = {"data": {"guardrailsResult": {"Allowed": True, "Reason": ""}}}
return mock
def _mock_blocked_response(reason="Prompt injection detected"):
mock = MagicMock(spec=httpx.Response)
mock.status_code = 200
mock.json.return_value = {"data": {"guardrailsResult": {"Allowed": False, "Reason": reason}}}
return mock
# ---------------------------------------------------------------------------
# Initialization tests
# ---------------------------------------------------------------------------
def test_init_requires_akto_base_url():
with patch.dict(os.environ, {}, clear=True):
with pytest.raises(ValueError, match="akto_base_url is required"):
AktoGuardrail(
akto_base_url="",
akto_api_key="test-token",
guardrail_name="test",
event_hook="pre_call",
)
def test_init_requires_api_key():
with patch.dict(os.environ, {}, clear=True):
with pytest.raises(ValueError, match="akto_api_key is required"):
AktoGuardrail(
akto_base_url="http://localhost:9090",
akto_api_key="",
guardrail_name="test",
event_hook="pre_call",
)
def test_init_from_env():
with patch.dict(
os.environ,
{
"AKTO_GUARDRAIL_API_BASE": "http://env-host:9090",
"AKTO_API_KEY": "env-token",
"AKTO_ACCOUNT_ID": "2000000",
"AKTO_VXLAN_ID": "42",
},
):
g = AktoGuardrail(guardrail_name="env-test", event_hook="post_call")
assert g.akto_base_url == "http://env-host:9090"
assert g.akto_api_key == "env-token"
assert g.guardrail_timeout == 5
assert g.akto_account_id == "2000000"
assert g.akto_vxlan_id == "42"
def test_init_defaults():
g = AktoGuardrail(
akto_base_url="http://localhost:9090",
akto_api_key="test-token",
guardrail_name="default-test",
event_hook="pre_call",
)
assert g.unreachable_fallback == "fail_closed"
assert g.guardrail_timeout == 5
assert g.akto_account_id == "1000000"
assert g.akto_vxlan_id == "0"
def test_background_tasks_per_instance():
a = AktoGuardrail(
akto_base_url="http://localhost:9090",
akto_api_key="test-token",
guardrail_name="instance-a",
event_hook="pre_call",
)
b = AktoGuardrail(
akto_base_url="http://localhost:9090",
akto_api_key="test-token",
guardrail_name="instance-b",
event_hook="post_call",
)
assert a.background_tasks is not b.background_tasks
# ---------------------------------------------------------------------------
# Payload format tests
# ---------------------------------------------------------------------------
def test_build_akto_payload_format(akto_validate, sample_inputs, sample_request_data):
payload = akto_validate.build_akto_payload(sample_inputs, sample_request_data, include_response=False)
assert payload["path"] == "/v1/chat/completions"
assert payload["method"] == "POST"
assert payload["type"] == "HTTP/1.1"
assert payload["akto_account_id"] == "1000000"
assert payload["akto_vxlan_id"] == "0"
assert payload["is_pending"] == "false"
assert payload["source"] == "MIRRORING"
assert payload["contextSource"] == "AGENTIC"
assert payload["ip"] == "10.0.0.1"
req_headers = json.loads(payload["requestHeaders"])
assert "content-type" in req_headers
req_wrapper = json.loads(payload["requestPayload"])
req_body = json.loads(req_wrapper["body"])
assert req_body["model"] == "gpt-4"
assert req_body["messages"][0]["content"] == "Hello, how are you?"
tag = json.loads(payload["tag"])
assert tag["gen-ai"] == "Gen AI"
assert payload["responsePayload"] == json.dumps({})
assert payload["time"].isdigit()
assert len(payload["time"]) >= 13
def test_build_akto_payload_with_response(akto_validate, sample_inputs, sample_request_data):
payload = akto_validate.build_akto_payload(sample_inputs, sample_request_data, include_response=True)
resp_wrapper = json.loads(payload["responsePayload"])
resp_body = json.loads(resp_wrapper["body"])
assert "choices" in resp_body
def test_build_akto_payload_custom_account_ids(sample_inputs, sample_request_data):
g = AktoGuardrail(
akto_base_url="http://localhost:9090",
akto_api_key="test-token",
akto_account_id="9999",
akto_vxlan_id="7",
guardrail_name="custom-ids-test",
event_hook="pre_call",
)
payload = g.build_akto_payload(sample_inputs, sample_request_data, include_response=False)
assert payload["akto_account_id"] == "9999"
assert payload["akto_vxlan_id"] == "7"
def test_build_query_params():
params = AktoGuardrail.build_query_params(guardrails=True, ingest_data=False)
assert params == {"akto_connector": "litellm", "guardrails": "true"}
params = AktoGuardrail.build_query_params(guardrails=False, ingest_data=True)
assert params == {"akto_connector": "litellm", "ingest_data": "true"}
params = AktoGuardrail.build_query_params(guardrails=True, ingest_data=True)
assert params == {
"akto_connector": "litellm",
"guardrails": "true",
"ingest_data": "true",
}
# ---------------------------------------------------------------------------
# Guardrail response handling
# ---------------------------------------------------------------------------
def test_handle_guardrail_response_allowed():
mock_resp = MagicMock(spec=httpx.Response)
mock_resp.status_code = 200
mock_resp.json.return_value = {"data": {"guardrailsResult": {"Allowed": True, "Reason": ""}}}
allowed, reason = AktoGuardrail.handle_guardrail_response(mock_resp)
assert allowed is True
assert reason == ""
def test_handle_guardrail_response_blocked():
mock_resp = MagicMock(spec=httpx.Response)
mock_resp.status_code = 200
mock_resp.json.return_value = {"data": {"guardrailsResult": {"Allowed": False, "Reason": "PII detected"}}}
allowed, reason = AktoGuardrail.handle_guardrail_response(mock_resp)
assert allowed is False
assert reason == "PII detected"
def test_handle_guardrail_response_missing_result():
mock_resp = MagicMock(spec=httpx.Response)
mock_resp.status_code = 200
mock_resp.json.return_value = {}
allowed, _ = AktoGuardrail.handle_guardrail_response(mock_resp)
assert allowed is True
def test_handle_guardrail_response_data_none():
mock_resp = MagicMock(spec=httpx.Response)
mock_resp.status_code = 200
mock_resp.json.return_value = {"data": None}
allowed, reason = AktoGuardrail.handle_guardrail_response(mock_resp)
assert allowed is True
assert reason == ""
def test_handle_guardrail_response_guardrails_result_not_dict():
mock_resp = MagicMock(spec=httpx.Response)
mock_resp.status_code = 200
mock_resp.json.return_value = {"data": {"guardrailsResult": "invalid"}}
allowed, reason = AktoGuardrail.handle_guardrail_response(mock_resp)
assert allowed is True
assert reason == ""
def test_handle_guardrail_response_non_dict():
mock_resp = MagicMock(spec=httpx.Response)
mock_resp.status_code = 200
mock_resp.json.return_value = "invalid"
allowed, _ = AktoGuardrail.handle_guardrail_response(mock_resp)
assert allowed is True
def test_handle_guardrail_response_error_status():
mock_resp = MagicMock(spec=httpx.Response)
mock_resp.status_code = 500
mock_resp.request = MagicMock()
with pytest.raises(httpx.HTTPStatusError):
AktoGuardrail.handle_guardrail_response(mock_resp)
def test_handle_guardrail_response_non_json_body():
mock_resp = MagicMock(spec=httpx.Response)
mock_resp.status_code = 200
mock_resp.request = MagicMock()
mock_resp.text = "<html>not json</html>"
mock_resp.json.side_effect = json.JSONDecodeError("Expecting value", "<html>", 0)
with pytest.raises(httpx.RequestError):
AktoGuardrail.handle_guardrail_response(mock_resp)
# ---------------------------------------------------------------------------
# Pre-call (akto-validate) — allowed
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_pre_call_allowed(akto_validate, sample_inputs, sample_request_data):
akto_validate.async_handler.post = AsyncMock(return_value=_mock_allowed_response())
result = await akto_validate.apply_guardrail(
inputs=sample_inputs,
request_data=sample_request_data,
input_type="request",
)
assert result == sample_inputs
akto_validate.async_handler.post.assert_called_once()
call_params = akto_validate.async_handler.post.call_args.kwargs["params"]
assert call_params.get("guardrails") == "true"
assert "ingest_data" not in call_params
# ---------------------------------------------------------------------------
# Pre-call (akto-validate) — blocked
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_pre_call_blocked(akto_validate, sample_inputs, sample_request_data):
akto_validate.async_handler.post = AsyncMock(
side_effect=[
_mock_blocked_response("PII detected"),
_mock_allowed_response(),
]
)
with pytest.raises(HTTPException) as exc_info:
await akto_validate.apply_guardrail(
inputs=sample_inputs,
request_data=sample_request_data,
input_type="request",
)
await asyncio.sleep(0)
await asyncio.sleep(0)
assert exc_info.value.status_code == 403
assert akto_validate.async_handler.post.call_count == 2
first_call_params = akto_validate.async_handler.post.call_args_list[0].kwargs["params"]
assert first_call_params.get("guardrails") == "true"
second_call_params = akto_validate.async_handler.post.call_args_list[1].kwargs["params"]
assert second_call_params.get("ingest_data") == "true"
assert "guardrails" not in second_call_params
second_payload = json.loads(akto_validate.async_handler.post.call_args_list[1].kwargs["data"])
assert second_payload["statusCode"] == "403"
resp_body = json.loads(second_payload["responsePayload"])
inner = json.loads(resp_body["body"])
assert inner["x-blocked-by"] == "Akto Proxy"
assert inner["reason"] == "PII detected"
# ---------------------------------------------------------------------------
# Pre-call (akto-validate) — response input is no-op
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_validate_response_noop(akto_validate, sample_inputs, sample_request_data):
akto_validate.async_handler.post = AsyncMock()
result = await akto_validate.apply_guardrail(
inputs=sample_inputs,
request_data=sample_request_data,
input_type="response",
)
assert result == sample_inputs
akto_validate.async_handler.post.assert_not_called()
# ---------------------------------------------------------------------------
# Post-call (akto-ingest) — combined guardrail + ingest
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_post_call_combined(akto_ingest, sample_inputs, sample_request_data):
akto_ingest.async_handler.post = AsyncMock(return_value=_mock_allowed_response())
result = await akto_ingest.apply_guardrail(
inputs=sample_inputs,
request_data=sample_request_data,
input_type="response",
)
await asyncio.sleep(0)
await asyncio.sleep(0)
assert result == sample_inputs
akto_ingest.async_handler.post.assert_called_once()
call_params = akto_ingest.async_handler.post.call_args.kwargs["params"]
assert call_params.get("guardrails") == "true"
assert call_params.get("ingest_data") == "true"
# ---------------------------------------------------------------------------
# Post-call (akto-ingest) — request input is no-op
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_ingest_request_noop(akto_ingest, sample_inputs, sample_request_data):
akto_ingest.async_handler.post = AsyncMock()
result = await akto_ingest.apply_guardrail(
inputs=sample_inputs,
request_data=sample_request_data,
input_type="request",
)
assert result == sample_inputs
akto_ingest.async_handler.post.assert_not_called()
# ---------------------------------------------------------------------------
# Fail-open / fail-closed
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_fail_open_on_unreachable():
g = AktoGuardrail(
akto_base_url="http://localhost:9090",
akto_api_key="test-token",
unreachable_fallback="fail_open",
guardrail_name="fail-open-test",
event_hook="pre_call",
)
g.async_handler.post = AsyncMock(side_effect=httpx.ConnectError("Connection refused"))
inputs = GenericGuardrailAPIInputs(texts=["test"], model="gpt-4")
result = await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request")
assert result.get("texts") == ["test"]
@pytest.mark.asyncio
async def test_fail_closed_on_unreachable():
g = AktoGuardrail(
akto_base_url="http://localhost:9090",
akto_api_key="test-token",
unreachable_fallback="fail_closed",
guardrail_name="fail-closed-test",
event_hook="pre_call",
)
g.async_handler.post = AsyncMock(side_effect=httpx.ConnectError("Connection refused"))
inputs = GenericGuardrailAPIInputs(texts=["test"], model="gpt-4")
with pytest.raises(HTTPException) as exc_info:
await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request")
assert exc_info.value.status_code == 503
def test_fail_closed_generic_message():
g = AktoGuardrail(
akto_base_url="http://localhost:9090",
akto_api_key="test-token",
unreachable_fallback="fail_closed",
guardrail_name="msg-test",
event_hook="pre_call",
)
with pytest.raises(HTTPException) as exc_info:
g.handle_unreachable(
inputs=GenericGuardrailAPIInputs(texts=["test"], model="gpt-4"),
error=Exception("http://internal-host:9090/secret-path"),
)
assert "internal-host" not in exc_info.value.detail
assert exc_info.value.detail == "Akto guardrail service unreachable"
# ---------------------------------------------------------------------------
# Helper method tests
# ---------------------------------------------------------------------------
def test_extract_request_path_from_metadata():
path = AktoGuardrail.extract_request_path({"metadata": {"user_api_key_request_route": "/v1/embeddings"}})
assert path == "/v1/embeddings"
def test_extract_request_path_fallback():
path = AktoGuardrail.extract_request_path({})
assert path == "/v1/chat/completions"
def test_extract_request_path_non_dict_metadata():
path = AktoGuardrail.extract_request_path({"metadata": "invalid"})
assert path == "/v1/chat/completions"
def test_resolve_metadata_value():
assert (
AktoGuardrail.resolve_metadata_value({"metadata": {"user_api_key_user_id": "u1"}}, "user_api_key_user_id")
== "u1"
)
assert (
AktoGuardrail.resolve_metadata_value(
{"litellm_metadata": {"user_api_key_team_id": "t1"}},
"user_api_key_team_id",
)
== "t1"
)
assert AktoGuardrail.resolve_metadata_value({}, "some_key") is None
assert AktoGuardrail.resolve_metadata_value(None, "some_key") is None
def test_resolve_metadata_value_non_dict_containers():
assert (
AktoGuardrail.resolve_metadata_value(
{"metadata": "invalid", "litellm_metadata": ["bad"]},
"some_key",
)
is None
)
def test_build_tag_metadata(akto_validate, sample_request_data):
tag = akto_validate.build_tag_metadata(sample_request_data)
assert tag["gen-ai"] == "Gen AI"
assert tag["user_id"] == "user-1"
assert tag["team_id"] == "team-1"

View file

@ -161,6 +161,24 @@ def test_document_inlining_example(disable_add_transform_inline_image_block):
"vision-gpt",
"http://example.com/image.png",
),
# data: URLs must never have #transform=inline appended — doing so
# corrupts the base64 payload (fixes #23583).
# URI schemes are case-insensitive (RFC 3986) so check all variants.
(
{"image_url": "data:image/png;base64,iVBORw0KGgo="},
"gpt-4",
"data:image/png;base64,iVBORw0KGgo=",
),
(
{"image_url": {"url": "data:image/jpeg;base64,/9j/4AAQ=="}},
"gpt-4",
{"url": "data:image/jpeg;base64,/9j/4AAQ=="},
),
(
{"image_url": "Data:image/png;base64,iVBORw0KGgo="},
"gpt-4",
"Data:image/png;base64,iVBORw0KGgo=",
),
],
)
def test_transform_inline(content, model, expected_url):

View file

@ -9,7 +9,9 @@ from unittest.mock import ANY, MagicMock, Mock, patch
import httpx
import pytest
sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system-path
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system-path
import litellm
@ -117,7 +119,9 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_imag
function_call_output = item
break
assert function_call_output is not None, "function_call_output not found in response"
assert (
function_call_output is not None
), "function_call_output not found in response"
assert function_call_output["call_id"] == "call_abc123"
# Check that the output is correctly transformed
@ -127,8 +131,12 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_imag
image_item = output[0]
# Should be transformed to Responses API format
assert image_item["type"] == "input_image", f"Expected type 'input_image', got '{image_item.get('type')}'"
assert image_item["image_url"] == test_image_base64, "image_url should be a flat string, not a nested object"
assert (
image_item["type"] == "input_image"
), f"Expected type 'input_image', got '{image_item.get('type')}'"
assert (
image_item["image_url"] == test_image_base64
), "image_url should be a flat string, not a nested object"
assert "detail" in image_item, "detail field should be present"
print("✓ Tool result with image correctly transformed to Responses API format")
@ -190,7 +198,9 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_text
function_call_output = item
break
assert function_call_output is not None, "function_call_output not found in response"
assert (
function_call_output is not None
), "function_call_output not found in response"
assert function_call_output["call_id"] == "call_abc123"
# Check that the output is correctly transformed to use input_text, not output_text
@ -200,12 +210,16 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_text
text_item = output[0]
# Should be transformed to use input_text for tool results in Responses API format
assert text_item["type"] == "input_text", (
f"Expected type 'input_text' for tool result, got '{text_item.get('type')}'"
)
assert text_item["text"] == "15 degrees", f"Expected text '15 degrees', got '{text_item.get('text')}'"
assert (
text_item["type"] == "input_text"
), f"Expected type 'input_text' for tool result, got '{text_item.get('type')}'"
assert (
text_item["text"] == "15 degrees"
), f"Expected text '15 degrees', got '{text_item.get('text')}'"
print("✓ Tool result with text correctly transformed to use input_text for Responses API format")
print(
"✓ Tool result with text correctly transformed to use input_text for Responses API format"
)
def test_openai_responses_chunk_parser_reasoning_summary():
@ -214,7 +228,9 @@ def test_openai_responses_chunk_parser_reasoning_summary():
)
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
chunk = {
"delta": "**Compar",
@ -246,7 +262,9 @@ def test_chunk_parser_string_output_text_delta_produces_text():
)
from litellm.types.utils import ModelResponseStream
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
chunk = {"type": "response.output_text.delta", "delta": "literal text"}
@ -267,7 +285,9 @@ def test_chunk_parser_enum_output_text_delta_produces_text():
from litellm.types.llms.openai import ResponsesAPIStreamEvents
from litellm.types.utils import ModelResponseStream
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
chunk = {"type": ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, "delta": "enum text"}
@ -288,7 +308,9 @@ def test_chunk_parser_function_call_added_produces_tool_use():
from litellm.types.llms.openai import ResponsesAPIStreamEvents
from litellm.types.utils import ModelResponseStream
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
chunk = {
"type": ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
@ -373,7 +395,9 @@ Tomorrow will bring its petitions and promises,
but for now the city breathes slow and wide,
and I learn to carry this small calm home."""
output_text = ResponseOutputText(annotations=[], text=poem_text, type="output_text", logprobs=[])
output_text = ResponseOutputText(
annotations=[], text=poem_text, type="output_text", logprobs=[]
)
output_message = ResponseOutputMessage(
id="msg_04c8021b8b3188a00068e9ae0b92f4819dac64d85b4abb67ec",
content=[output_text],
@ -385,7 +409,9 @@ and I learn to carry this small calm home."""
# Create usage information
usage = ResponseAPIUsage(
input_tokens=16,
input_tokens_details=InputTokensDetails(audio_tokens=None, cached_tokens=0, text_tokens=None),
input_tokens_details=InputTokensDetails(
audio_tokens=None, cached_tokens=0, text_tokens=None
),
output_tokens=195,
output_tokens_details=OutputTokensDetails(reasoning_tokens=0, text_tokens=None),
total_tokens=211,
@ -597,7 +623,9 @@ def test_transform_request_single_char_keys_not_matched():
assert result_correct.get("metadata") == {"user_id": "123"}
assert result_correct.get("previous_response_id") == "resp_abc"
print("✓ Single-character keys are not incorrectly matched to metadata/previous_response_id")
print(
"✓ Single-character keys are not incorrectly matched to metadata/previous_response_id"
)
# =============================================================================
@ -617,7 +645,9 @@ def test_message_done_does_not_emit_is_finished():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
chunk = {
"type": "response.output_item.done",
@ -629,9 +659,9 @@ def test_message_done_does_not_emit_is_finished():
# After the fix, message completion should NOT set finish_reason
# ModelResponseStream doesn't have is_finished - check finish_reason instead
assert len(result.choices) > 0, "result should have choices"
assert result.choices[0].finish_reason is None or result.choices[0].finish_reason == "", (
"message completion should not emit finish_reason"
)
assert (
result.choices[0].finish_reason is None or result.choices[0].finish_reason == ""
), "message completion should not emit finish_reason"
def test_response_completed_emits_is_finished():
@ -643,7 +673,9 @@ def test_response_completed_emits_is_finished():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
chunk = {"type": "response.completed"}
@ -651,7 +683,9 @@ def test_response_completed_emits_is_finished():
# response.completed should emit finish_reason='stop'
assert len(result.choices) > 0, "result should have choices"
assert result.choices[0].finish_reason == "stop", "response.completed should emit finish_reason='stop'"
assert (
result.choices[0].finish_reason == "stop"
), "response.completed should emit finish_reason='stop'"
def test_response_completed_with_function_calls_emits_tool_calls_finish_reason():
@ -670,7 +704,9 @@ def test_response_completed_with_function_calls_emits_tool_calls_finish_reason()
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
# Simulate a response.completed event with function_call in output
# This matches what Azure/OpenAI sends for gpt-5.1-codex-mini and similar models
@ -696,9 +732,9 @@ def test_response_completed_with_function_calls_emits_tool_calls_finish_reason()
# response.completed with function_call should emit finish_reason='tool_calls'
assert len(result.choices) > 0, "result should have choices"
assert result.choices[0].finish_reason == "tool_calls", (
"response.completed with function_call output should emit finish_reason='tool_calls'"
)
assert (
result.choices[0].finish_reason == "tool_calls"
), "response.completed with function_call output should emit finish_reason='tool_calls'"
def test_response_completed_with_message_only_emits_stop_finish_reason():
@ -709,7 +745,9 @@ def test_response_completed_with_message_only_emits_stop_finish_reason():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
# Simulate a response.completed event with only message output
chunk = {
@ -733,10 +771,9 @@ def test_response_completed_with_message_only_emits_stop_finish_reason():
# response.completed with only message should emit finish_reason='stop'
assert len(result.choices) > 0, "result should have choices"
assert result.choices[0].finish_reason == "stop", (
"response.completed with only message output should emit finish_reason='stop'"
)
assert (
result.choices[0].finish_reason == "stop"
), "response.completed with only message output should emit finish_reason='stop'"
def test_response_completed_preserves_usage_with_cached_tokens():
@ -752,7 +789,9 @@ def test_response_completed_preserves_usage_with_cached_tokens():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
chunk = {
"type": "response.completed",
@ -781,12 +820,18 @@ def test_response_completed_preserves_usage_with_cached_tokens():
result = iterator.chunk_parser(chunk)
assert result.usage is not None, "usage should be set on response.completed chunk"
assert result.usage.prompt_tokens == 1226, "prompt_tokens should map from input_tokens"
assert result.usage.completion_tokens == 5, "completion_tokens should map from output_tokens"
assert result.usage.prompt_tokens_details is not None, "prompt_tokens_details should be set"
assert result.usage.prompt_tokens_details.cached_tokens == 1024, (
"cached_tokens should be preserved from input_tokens_details"
)
assert (
result.usage.prompt_tokens == 1226
), "prompt_tokens should map from input_tokens"
assert (
result.usage.completion_tokens == 5
), "completion_tokens should map from output_tokens"
assert (
result.usage.prompt_tokens_details is not None
), "prompt_tokens_details should be set"
assert (
result.usage.prompt_tokens_details.cached_tokens == 1024
), "cached_tokens should be preserved from input_tokens_details"
def test_function_call_done_emits_is_finished():
@ -800,7 +845,9 @@ def test_function_call_done_emits_is_finished():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
chunk = {
"type": "response.output_item.done",
@ -820,9 +867,9 @@ def test_function_call_done_emits_is_finished():
"output_item.done for function_call must not emit finish_reason; "
"response.completed is responsible for the terminal finish_reason"
)
assert not result.choices[0].delta.tool_calls, (
"output_item.done for function_call must not include a duplicate tool_calls delta"
)
assert not result.choices[
0
].delta.tool_calls, "output_item.done for function_call must not include a duplicate tool_calls delta"
def test_text_plus_tool_calls_sequence():
@ -837,7 +884,9 @@ def test_text_plus_tool_calls_sequence():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
# Simulate the sequence from OpenAI Responses API
chunks = [
@ -876,23 +925,28 @@ def test_text_plus_tool_calls_sequence():
# Check message done (index 2) does NOT have finish_reason set
message_done_result = results[2]
assert len(message_done_result.choices) > 0, "message done should have choices"
assert message_done_result.choices[0].finish_reason is None or message_done_result.choices[0].finish_reason == "", (
"message done should not have finish_reason"
)
assert (
message_done_result.choices[0].finish_reason is None
or message_done_result.choices[0].finish_reason == ""
), "message done should not have finish_reason"
# Check function_call done (index 5) does NOT have finish_reason set
# (response.completed is responsible for the terminal finish_reason)
function_done_result = results[5]
assert len(function_done_result.choices) > 0, "function_call done should have choices"
assert function_done_result.choices[0].finish_reason is None, (
"output_item.done for function_call must not emit finish_reason"
)
assert (
len(function_done_result.choices) > 0
), "function_call done should have choices"
assert (
function_done_result.choices[0].finish_reason is None
), "output_item.done for function_call must not emit finish_reason"
# Check response.completed (index 6) has finish_reason='stop'
# (the mock chunk has no nested 'response' data, so has_function_calls is False → 'stop')
completed_result = results[6]
assert len(completed_result.choices) > 0, "response.completed should have choices"
assert completed_result.choices[0].finish_reason == "stop", "response.completed should have finish_reason='stop'"
assert (
completed_result.choices[0].finish_reason == "stop"
), "response.completed should have finish_reason='stop'"
# =============================================================================
@ -958,7 +1012,9 @@ def test_tool_message_output_uses_input_text_not_output_text():
output = function_call_output["output"]
assert isinstance(output, list), f"output should be a list, got {type(output)}"
assert len(output) == 1
assert output[0]["type"] == "input_text", f"Expected input_text, got {output[0].get('type')}"
assert (
output[0]["type"] == "input_text"
), f"Expected input_text, got {output[0].get('type')}"
assert output[0]["text"] == '{"temperature": 15, "condition": "sunny"}'
print("✓ Tool message output correctly uses input_text type")
@ -1144,9 +1200,13 @@ def test_map_reasoning_effort_adds_summary_detailed():
assert result is not None, f"Result should not be None for effort={effort}"
assert result["effort"] == effort, f"Effort should be {effort}"
assert "summary" not in result, f"Summary should NOT be present by default for effort={effort}"
assert (
"summary" not in result
), f"Summary should NOT be present by default for effort={effort}"
print(f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}' (no summary by default)")
print(
f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}' (no summary by default)"
)
# Test 2: With flag enabled - summary IS added
litellm.reasoning_auto_summary = True
@ -1156,9 +1216,9 @@ def test_map_reasoning_effort_adds_summary_detailed():
assert result is not None, f"Result should not be None for effort={effort}"
assert result["effort"] == effort, f"Effort should be {effort}"
assert result["summary"] == "detailed", (
f"Summary should be 'detailed' when flag is enabled for effort={effort}"
)
assert (
result["summary"] == "detailed"
), f"Summary should be 'detailed' when flag is enabled for effort={effort}"
print(
f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}', summary='detailed' (flag enabled)"
@ -1169,7 +1229,9 @@ def test_map_reasoning_effort_adds_summary_detailed():
os.environ["LITELLM_REASONING_AUTO_SUMMARY"] = "true"
result = handler._map_reasoning_effort("high")
assert result["summary"] == "detailed", "Summary should be 'detailed' when env var is enabled"
assert (
result["summary"] == "detailed"
), "Summary should be 'detailed' when env var is enabled"
print("✓ LITELLM_REASONING_AUTO_SUMMARY env var works correctly")
# Test 4: Dict input is passed through as-is (no modification)
@ -1188,7 +1250,9 @@ def test_map_reasoning_effort_adds_summary_detailed():
assert result_unknown is None
print("✓ Unknown reasoning_effort values return None")
print("✓ All reasoning_effort behaviors work correctly with flag/env var control")
print(
"✓ All reasoning_effort behaviors work correctly with flag/env var control"
)
finally:
# Restore original values
@ -1264,7 +1328,9 @@ def test_transform_response_preserves_annotations():
# Create usage information
usage = ResponseAPIUsage(
input_tokens=10,
input_tokens_details=InputTokensDetails(audio_tokens=None, cached_tokens=0, text_tokens=None),
input_tokens_details=InputTokensDetails(
audio_tokens=None, cached_tokens=0, text_tokens=None
),
output_tokens=20,
output_tokens_details=OutputTokensDetails(reasoning_tokens=0, text_tokens=None),
total_tokens=30,
@ -1351,9 +1417,13 @@ def test_transform_response_preserves_annotations():
assert choice.message.content == "Here is some information with citations."
# Check that annotations are preserved
assert hasattr(choice.message, "annotations"), "Message should have annotations attribute"
assert hasattr(
choice.message, "annotations"
), "Message should have annotations attribute"
assert choice.message.annotations is not None, "Annotations should not be None"
assert len(choice.message.annotations) == 2, f"Expected 2 annotations, got {len(choice.message.annotations)}"
assert (
len(choice.message.annotations) == 2
), f"Expected 2 annotations, got {len(choice.message.annotations)}"
# Verify annotation content
annotation1 = choice.message.annotations[0]
@ -1375,7 +1445,9 @@ def test_transform_response_preserves_annotations():
assert result.usage.completion_tokens == 20
assert result.usage.total_tokens == 30
print("✓ Annotations from Responses API are correctly preserved in Chat Completions format")
print(
"✓ Annotations from Responses API are correctly preserved in Chat Completions format"
)
def test_apply_patch_tool_call_converted_to_chat_completion_tool_call():
@ -1512,6 +1584,8 @@ def test_apply_patch_tool_call_converted_to_chat_completion_tool_call():
assert args["type"] == "create_file"
assert args["path"] == "hello.py"
assert "print('hello world')" in args["diff"]
def test_multi_tool_call_stream_no_premature_finish():
"""
Regression test for multi-tool-call streaming bug.
@ -1538,18 +1612,26 @@ def test_multi_tool_call_stream_no_premature_finish():
OpenAiResponsesToChatCompletionStreamIterator,
)
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
chunks = [
# 0: response created
{"type": "response.created", "response": {"id": "resp_001", "status": "in_progress"}},
{
"type": "response.created",
"response": {"id": "resp_001", "status": "in_progress"},
},
# 1: first tool call added
{
"type": "response.output_item.added",
"item": {"type": "function_call", "name": "read_file", "call_id": "call_1"},
},
# 2: first tool call arguments delta
{"type": "response.function_call_arguments.delta", "delta": '{"path":"/etc/hostname"}'},
{
"type": "response.function_call_arguments.delta",
"delta": '{"path":"/etc/hostname"}',
},
# 3: first tool call done ← must NOT emit finish_reason
{
"type": "response.output_item.done",
@ -1608,10 +1690,12 @@ def test_multi_tool_call_stream_no_premature_finish():
r = results[done_idx]
assert r is not None, f"{label}: chunk_parser must return a result"
assert len(r.choices) > 0, f"{label}: result must have choices"
assert r.choices[0].finish_reason is None, (
f"{label}: output_item.done must not emit finish_reason (stream would terminate prematurely)"
)
assert not r.choices[0].delta.tool_calls, (
assert (
r.choices[0].finish_reason is None
), f"{label}: output_item.done must not emit finish_reason (stream would terminate prematurely)"
assert not r.choices[
0
].delta.tool_calls, (
f"{label}: output_item.done must not include a duplicate tool_calls delta"
)
@ -1623,12 +1707,12 @@ def test_multi_tool_call_stream_no_premature_finish():
r = results[added_idx]
if r is not None and r.choices and r.choices[0].delta.tool_calls:
tc = r.choices[0].delta.tool_calls[0]
assert tc.function.name == expected_name, (
f"output_item.added for {expected_name}: tool_call name mismatch"
)
assert tc.id == expected_call_id, (
f"output_item.added for {expected_name}: call_id mismatch"
)
assert (
tc.function.name == expected_name
), f"output_item.added for {expected_name}: tool_call name mismatch"
assert (
tc.id == expected_call_id
), f"output_item.added for {expected_name}: call_id mismatch"
# 3. argument delta events (indices 2 and 5) should carry arguments
for delta_idx, expected_args, label in [
@ -1638,17 +1722,17 @@ def test_multi_tool_call_stream_no_premature_finish():
r = results[delta_idx]
if r is not None and r.choices and r.choices[0].delta.tool_calls:
tc = r.choices[0].delta.tool_calls[0]
assert tc.function.arguments == expected_args, (
f"{label}: argument delta mismatch"
)
assert (
tc.function.arguments == expected_args
), f"{label}: argument delta mismatch"
# 4. Only response.completed (index 7) emits the terminal finish_reason
completed_result = results[7]
assert completed_result is not None, "response.completed must return a result"
assert len(completed_result.choices) > 0, "response.completed must have choices"
assert completed_result.choices[0].finish_reason == "tool_calls", (
"response.completed with function_call outputs must emit finish_reason='tool_calls'"
)
assert (
completed_result.choices[0].finish_reason == "tool_calls"
), "response.completed with function_call outputs must emit finish_reason='tool_calls'"
# 5. No chunk before the last one should have finish_reason set
for idx, r in enumerate(results[:-1]):
@ -1658,7 +1742,9 @@ def test_multi_tool_call_stream_no_premature_finish():
f"— only response.completed should terminate the stream"
)
print("✓ Multi-tool-call stream completes without premature finish_reason termination")
print(
"✓ Multi-tool-call stream completes without premature finish_reason termination"
)
# =============================================================================
@ -1790,7 +1876,10 @@ def test_parallel_tool_calls_comprehensive_streaming_integration():
chunks = [
# 0: response.created
{"type": "response.created", "response": {"id": "resp_001", "status": "in_progress"}},
{
"type": "response.created",
"response": {"id": "resp_001", "status": "in_progress"},
},
# 1: call_1 (read_file) added — output_index=0
{
"type": "response.output_item.added",
@ -1873,7 +1962,9 @@ def test_parallel_tool_calls_comprehensive_streaming_integration():
},
]
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
results = [iterator.chunk_parser(chunk) for chunk in chunks]
# 1. output_item.done events (indices 4 and 8) must NOT emit finish_reason
@ -1885,7 +1976,9 @@ def test_parallel_tool_calls_comprehensive_streaming_integration():
f"{label}: output_item.done must not emit finish_reason "
f"(would prematurely terminate stream before subsequent tool calls arrive)"
)
assert not r.choices[0].delta.tool_calls, (
assert not r.choices[
0
].delta.tool_calls, (
f"{label}: output_item.done must not emit a duplicate tool_calls delta"
)
@ -1919,7 +2012,9 @@ def test_parallel_tool_calls_comprehensive_streaming_integration():
for tc in tool_calls:
if tc.function and tc.function.arguments:
idx = tc.index
assembled_args[idx] = assembled_args.get(idx, "") + tc.function.arguments
assembled_args[idx] = (
assembled_args.get(idx, "") + tc.function.arguments
)
# delta 1 = '{"path":' + delta 2 = '"/etc/foo"}' → '{"path":"/etc/foo"}'
assert assembled_args.get(0) == '{"path":"/etc/foo"}', (
@ -1938,16 +2033,16 @@ def test_parallel_tool_calls_comprehensive_streaming_integration():
for i, r in enumerate(results)
if r is not None and r.choices and r.choices[0].finish_reason
]
assert len(finish_events) == 1, (
f"Expected exactly 1 finish event, got {len(finish_events)}: {finish_events}"
)
assert (
len(finish_events) == 1
), f"Expected exactly 1 finish event, got {len(finish_events)}: {finish_events}"
assert finish_events[0][0] == len(chunks) - 1, (
f"Finish event must be at the last chunk (index {len(chunks) - 1}), "
f"but was at index {finish_events[0][0]}"
)
assert finish_events[0][1] == "tool_calls", (
f"Terminal finish_reason must be 'tool_calls', got '{finish_events[0][1]}'"
)
assert (
finish_events[0][1] == "tool_calls"
), f"Terminal finish_reason must be 'tool_calls', got '{finish_events[0][1]}'"
# 5. Parallel tool calls have distinct indices matching output_index (0 and 1)
# Collect indices from output_item.added chunks only (they carry the call id)
@ -1958,16 +2053,19 @@ def test_parallel_tool_calls_comprehensive_streaming_integration():
for tc in r.choices[0].delta.tool_calls
if tc.id # output_item.added chunks carry the id; argument deltas do not
]
assert set(added_tool_call_indices) == {0, 1}, (
f"Parallel tool calls must have distinct indices {{0, 1}}, got: {set(added_tool_call_indices)}"
)
assert set(added_tool_call_indices) == {
0,
1,
}, f"Parallel tool calls must have distinct indices {{0, 1}}, got: {set(added_tool_call_indices)}"
print("✓ Parallel tool calls with split argument deltas stream correctly end-to-end")
print(
"✓ Parallel tool calls with split argument deltas stream correctly end-to-end"
)
def test_map_optional_params_preserves_reasoning_summary():
"""Test that reasoning_effort dict with summary field is preserved.
Regression test for: User reported that summary field was being dropped
when routing to Responses API. The dict format should be fully preserved.
"""
@ -1992,6 +2090,97 @@ def test_map_optional_params_preserves_reasoning_summary():
# Verify reasoning_effort dict with summary was fully preserved
assert "reasoning" in responses_api_request
assert responses_api_request["reasoning"] == {"effort": "high", "summary": "detailed"}
assert responses_api_request["reasoning"] == {
"effort": "high",
"summary": "detailed",
}
assert responses_api_request["reasoning"]["effort"] == "high"
assert responses_api_request["reasoning"]["summary"] == "detailed"
def test_convert_chat_completion_file_type_to_input_file():
"""
Test that Chat Completion content with type 'file' is correctly mapped
to Responses API 'input_file' format, not stringified as 'input_text'.
Regression test for https://github.com/BerriAI/litellm/issues/23588
"""
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
)
handler = LiteLLMResponsesTransformationHandler()
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "What is in this PDF?"},
{
"type": "file",
"file": {
"file_data": "data:application/pdf;base64,JVBERi0xLjQK",
"filename": "test.pdf",
},
},
],
}
]
input_items, instructions = (
handler.convert_chat_completion_messages_to_responses_api(messages)
)
assert len(input_items) == 1
msg = input_items[0]
assert msg["type"] == "message"
assert msg["role"] == "user"
content = msg["content"]
assert len(content) == 2
# First item should be the text
assert content[0]["type"] == "input_text"
assert content[0]["text"] == "What is in this PDF?"
# Second item should be input_file, NOT input_text with stringified dict
assert content[1]["type"] == "input_file"
assert content[1]["file_data"] == "data:application/pdf;base64,JVBERi0xLjQK"
assert content[1]["filename"] == "test.pdf"
# Ensure it does NOT have the nested 'file' key
assert "file" not in content[1]
def test_convert_chat_completion_file_type_with_file_id():
"""
Test that Chat Completion content with type 'file' using file_id is correctly mapped.
"""
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
)
handler = LiteLLMResponsesTransformationHandler()
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "Summarize this file."},
{
"type": "file",
"file": {
"file_id": "file-abc123",
},
},
],
}
]
input_items, instructions = (
handler.convert_chat_completion_messages_to_responses_api(messages)
)
content = input_items[0]["content"]
assert content[1]["type"] == "input_file"
assert content[1]["file_id"] == "file-abc123"
assert "file_data" not in content[1]

View file

@ -16,13 +16,9 @@ class TestLangsmithLoggerInit:
Note: The current implementation has some edge cases in the sampling rate logic.
"""
@patch("asyncio.create_task")
@patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": "1"}, clear=False)
def test_langsmith_sampling_rate_parameter_respected_with_valid_env(
self, mock_create_task
):
def test_langsmith_sampling_rate_parameter_respected_with_valid_env(self):
"""Test that langsmith_sampling_rate parameter is properly set when env var condition is met."""
# When there's a valid integer in env var, the parameter should be used due to 'or' logic
sampling_rate = 0.5
logger = LangsmithLogger(
langsmith_api_key="test-key",
@ -30,58 +26,47 @@ class TestLangsmithLoggerInit:
langsmith_sampling_rate=sampling_rate,
)
# With the current 'or' logic and valid env var, the parameter should be used
assert (
logger.sampling_rate == sampling_rate
), f"Expected sampling_rate to be {sampling_rate}, got {logger.sampling_rate}"
@patch("asyncio.create_task")
@patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": "1"}, clear=False)
def test_langsmith_sampling_rate_zero_parameter_falls_back_to_env(
self, mock_create_task
):
def test_langsmith_sampling_rate_zero_parameter_falls_back_to_env(self):
"""Test that 0.0 parameter falls back to env var due to falsy value."""
# This demonstrates the current behavior where 0.0 is falsy and falls back to env
logger = LangsmithLogger(
langsmith_api_key="test-key",
langsmith_project="test-project",
langsmith_sampling_rate=0.0, # This is falsy!
langsmith_sampling_rate=0.0,
)
# Due to current 'or' logic, 0.0 falls back to env var
assert (
logger.sampling_rate == 1.0
), f"Expected sampling_rate to fall back to 1.0 from env, got {logger.sampling_rate}"
@patch("asyncio.create_task")
@patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": "1"}, clear=False)
def test_langsmith_sampling_rate_from_integer_env_var(self, mock_create_task):
def test_langsmith_sampling_rate_from_integer_env_var(self):
"""Test that sampling rate uses environment variable when parameter not provided and env var is integer."""
logger = LangsmithLogger(
langsmith_api_key="test-key", langsmith_project="test-project"
)
# Should use env var since it's a valid integer
assert (
logger.sampling_rate == 1.0
), f"Expected sampling_rate to be 1.0 from env var, got {logger.sampling_rate}"
@patch("asyncio.create_task")
@patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": "0.8"}, clear=False)
def test_langsmith_sampling_rate_decimal_env_var_ignored(self, mock_create_task):
def test_langsmith_sampling_rate_decimal_env_var_ignored(self):
"""Test that decimal environment variables are ignored due to isdigit() check."""
logger = LangsmithLogger(
langsmith_api_key="test-key", langsmith_project="test-project"
)
# Decimal env vars are ignored due to isdigit() check, falls back to 1.0
assert (
logger.sampling_rate == 1.0
), f"Expected sampling_rate to default to 1.0 (decimal env ignored), got {logger.sampling_rate}"
@patch("asyncio.create_task")
@patch.dict(os.environ, {}, clear=True)
def test_langsmith_sampling_rate_default_value(self, mock_create_task):
def test_langsmith_sampling_rate_default_value(self):
"""Test that sampling rate defaults to 1.0 when no parameter or env var provided."""
logger = LangsmithLogger(
langsmith_api_key="test-key", langsmith_project="test-project"
@ -91,9 +76,8 @@ class TestLangsmithLoggerInit:
logger.sampling_rate == 1.0
), f"Expected default sampling_rate to be 1.0, got {logger.sampling_rate}"
@patch("asyncio.create_task")
@patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": "invalid"}, clear=False)
def test_langsmith_sampling_rate_invalid_env_var_defaults(self, mock_create_task):
def test_langsmith_sampling_rate_invalid_env_var_defaults(self):
"""Test that invalid environment variable falls back to default value."""
logger = LangsmithLogger(
langsmith_api_key="test-key", langsmith_project="test-project"
@ -103,9 +87,8 @@ class TestLangsmithLoggerInit:
logger.sampling_rate == 1.0
), f"Expected sampling_rate to default to 1.0 with invalid env var, got {logger.sampling_rate}"
@patch("asyncio.create_task")
@patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": ""}, clear=False)
def test_langsmith_sampling_rate_empty_env_var_defaults(self, mock_create_task):
def test_langsmith_sampling_rate_empty_env_var_defaults(self):
"""Test that empty environment variable falls back to default value."""
logger = LangsmithLogger(
langsmith_api_key="test-key", langsmith_project="test-project"
@ -115,14 +98,12 @@ class TestLangsmithLoggerInit:
logger.sampling_rate == 1.0
), f"Expected sampling_rate to default to 1.0 with empty env var, got {logger.sampling_rate}"
@patch("asyncio.create_task")
def test_langsmith_sampling_rate_attribute_exists(self, mock_create_task):
def test_langsmith_sampling_rate_attribute_exists(self):
"""Test that the sampling_rate attribute is always set on the logger instance."""
logger = LangsmithLogger(
langsmith_api_key="test-key", langsmith_project="test-project"
)
# Verify the attribute exists and is a float
assert hasattr(
logger, "sampling_rate"
), "LangsmithLogger should have sampling_rate attribute"
@ -133,6 +114,93 @@ class TestLangsmithLoggerInit:
logger.sampling_rate >= 0.0
), f"sampling_rate should be non-negative, got {logger.sampling_rate}"
@patch.object(LangsmithLogger, "_start_periodic_flush_task", return_value=None)
def test_langsmith_init_skips_periodic_flush_without_running_loop(
self, mock_start_periodic_flush_task
):
"""Test that sync initialization leaves the periodic flush task unset."""
logger = LangsmithLogger(
langsmith_api_key="test-key", langsmith_project="test-project"
)
assert logger is not None
mock_start_periodic_flush_task.assert_called_once()
assert logger._flush_task is None
@patch("asyncio.get_running_loop", side_effect=RuntimeError("no running event loop"))
def test_start_periodic_flush_task_returns_none_without_running_loop(
self, mock_get_running_loop
):
"""Test that helper returns None when no running event loop exists."""
with patch.object(LangsmithLogger, "_start_periodic_flush_task", return_value=None):
logger = LangsmithLogger(
langsmith_api_key="test-key",
langsmith_project="test-project",
)
mock_get_running_loop.reset_mock()
assert logger._start_periodic_flush_task() is None
mock_get_running_loop.assert_called_once()
@patch("asyncio.get_running_loop")
def test_langsmith_init_starts_periodic_flush_with_running_loop(
self, mock_get_running_loop
):
"""Test that init schedules periodic flush when a running loop exists."""
mock_loop = MagicMock()
mock_task = MagicMock()
mock_loop.create_task.return_value = mock_task
mock_get_running_loop.return_value = mock_loop
logger = LangsmithLogger(
langsmith_api_key="test-key", langsmith_project="test-project"
)
assert logger._flush_task == mock_task
mock_loop.create_task.assert_called_once()
scheduled_coro = mock_loop.create_task.call_args.args[0]
scheduled_coro.close()
@pytest.mark.asyncio
async def test_async_log_success_event_lazily_starts_periodic_flush(self):
"""Test that async logging lazily starts periodic flush after sync init."""
with patch.object(LangsmithLogger, "_start_periodic_flush_task", return_value=None):
logger = LangsmithLogger(
langsmith_api_key="test-key",
langsmith_project="test-project",
)
logger._get_sampling_rate_to_use_for_request = MagicMock(return_value=1.0)
logger._get_credentials_to_use_for_request = MagicMock(
return_value=logger.default_credentials
)
logger._prepare_log_data = MagicMock(return_value={"id": "run-id"})
logger._start_periodic_flush_task = MagicMock(return_value=MagicMock())
await logger.async_log_success_event({}, {}, None, None)
logger._start_periodic_flush_task.assert_called_once()
assert len(logger.log_queue) == 1
@pytest.mark.asyncio
async def test_async_log_failure_event_lazily_starts_periodic_flush(self):
"""Test that async failure logging lazily starts periodic flush after sync init."""
with patch.object(LangsmithLogger, "_start_periodic_flush_task", return_value=None):
logger = LangsmithLogger(
langsmith_api_key="test-key",
langsmith_project="test-project",
)
logger._get_sampling_rate_to_use_for_request = MagicMock(return_value=1.0)
logger._get_credentials_to_use_for_request = MagicMock(
return_value=logger.default_credentials
)
logger._prepare_log_data = MagicMock(return_value={"id": "run-id"})
logger._start_periodic_flush_task = MagicMock(return_value=MagicMock())
await logger.async_log_failure_event({}, {}, None, None)
logger._start_periodic_flush_task.assert_called_once()
assert len(logger.log_queue) == 1
class TestLangsmithPrepareLogData:
"""Regression test for #24001: _prepare_log_data must inject

View file

@ -1,4 +1,5 @@
import base64
import json
from unittest.mock import MagicMock, patch
import pytest
@ -10,9 +11,11 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
BedrockImageProcessor,
anthropic_messages_pt,
_convert_to_bedrock_tool_call_invoke,
convert_to_gemini_tool_call_result,
ollama_pt,
sanitize_messages_for_tool_calling,
)
from litellm.types.llms.openai import ChatCompletionToolMessage
def test_ollama_pt_simple_messages():
@ -551,6 +554,175 @@ def test_convert_gemini_tool_call_result_with_image_url():
assert isinstance(result2, list) and any("inline_data" in p for p in result2)
def test_convert_gemini_tool_call_result_with_anthropic_image_block():
"""
Test that Anthropic-native image blocks in tool_result list content are
converted to Gemini inline_data instead of being silently dropped.
Fixes: https://github.com/BerriAI/litellm/issues/23712
"""
tiny_png_b64 = base64.b64encode(b"PNG_PLACEHOLDER").decode()
message = ChatCompletionToolMessage(
role="tool",
tool_call_id="call_123",
content=[
{
"type": "image",
"source": {
"type": "base64",
"media_type": "image/png",
"data": tiny_png_b64,
},
}
],
)
last_message_with_tool_calls = {
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": "call_123",
"type": "function",
"index": 0,
"function": {"name": "read_file", "arguments": "{}"},
}
],
}
result = convert_to_gemini_tool_call_result(
message=message,
last_message_with_tool_calls=last_message_with_tool_calls,
)
assert isinstance(result, list), "expected a list of parts"
inline_parts = [p for p in result if "inline_data" in p]
assert len(inline_parts) == 1, "expected exactly one inline_data part"
assert inline_parts[0]["inline_data"]["mime_type"] == "image/png"
assert inline_parts[0]["inline_data"]["data"] == tiny_png_b64
def test_convert_gemini_tool_call_result_with_multiple_anthropic_image_blocks():
"""
Test that multiple Anthropic-native image blocks in a single tool_result
are all preserved as separate inline_data parts instead of only the last
one being kept.
Fixes: https://github.com/BerriAI/litellm/issues/23712
"""
png_b64 = base64.b64encode(b"PNG_PLACEHOLDER").decode()
jpeg_b64 = base64.b64encode(b"JPEG_PLACEHOLDER").decode()
message = ChatCompletionToolMessage(
role="tool",
tool_call_id="call_multi",
content=[
{"type": "text", "text": "here are two images"},
{
"type": "image",
"source": {"type": "base64", "media_type": "image/png", "data": png_b64},
},
{
"type": "image",
"source": {"type": "base64", "media_type": "image/jpeg", "data": jpeg_b64},
},
],
)
last_message_with_tool_calls = {
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": "call_multi",
"type": "function",
"index": 0,
"function": {"name": "screenshot", "arguments": "{}"},
}
],
}
result = convert_to_gemini_tool_call_result(
message=message,
last_message_with_tool_calls=last_message_with_tool_calls,
)
assert isinstance(result, list), "expected a list of parts"
inline_parts = [p for p in result if "inline_data" in p]
assert len(inline_parts) == 2, f"expected 2 inline_data parts, got {len(inline_parts)}"
mime_types = {p["inline_data"]["mime_type"] for p in inline_parts}
assert mime_types == {"image/png", "image/jpeg"}
def test_convert_gemini_tool_call_result_with_data_url_string():
"""
Test that a data-URL string in tool_result content is converted to
Gemini inline_data instead of being passed as plain text.
Fixes: https://github.com/BerriAI/litellm/issues/23712
"""
tiny_png_b64 = base64.b64encode(b"PNG_PLACEHOLDER").decode()
message = ChatCompletionToolMessage(
role="tool",
tool_call_id="call_456",
content=f"data:image/png;base64,{tiny_png_b64}",
)
last_message_with_tool_calls = {
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": "call_456",
"type": "function",
"index": 0,
"function": {"name": "read_file", "arguments": "{}"},
}
],
}
result = convert_to_gemini_tool_call_result(
message=message,
last_message_with_tool_calls=last_message_with_tool_calls,
)
assert isinstance(result, list), "expected a list of parts"
inline_parts = [p for p in result if "inline_data" in p]
assert len(inline_parts) == 1, "data-URL image string was not converted to inline_data"
assert inline_parts[0]["inline_data"]["mime_type"] == "image/png"
assert inline_parts[0]["inline_data"]["data"] == tiny_png_b64
def test_convert_gemini_tool_call_result_with_data_url_extra_params():
"""
Test that a data-URL with extra MIME parameters (e.g. charset) produces
a clean mime_type without the extra parameters.
"""
tiny_png_b64 = base64.b64encode(b"PNG_PLACEHOLDER").decode()
message = ChatCompletionToolMessage(
role="tool",
tool_call_id="call_extra",
content=f"data:image/png;charset=UTF-8;base64,{tiny_png_b64}",
)
last_message_with_tool_calls = {
"role": "assistant",
"content": "",
"tool_calls": [
{
"id": "call_extra",
"type": "function",
"index": 0,
"function": {"name": "read_file", "arguments": "{}"},
}
],
}
result = convert_to_gemini_tool_call_result(
message=message,
last_message_with_tool_calls=last_message_with_tool_calls,
)
assert isinstance(result, list), "expected a list of parts"
inline_parts = [p for p in result if "inline_data" in p]
assert len(inline_parts) == 1
assert inline_parts[0]["inline_data"]["mime_type"] == "image/png", (
f"expected clean 'image/png', got '{inline_parts[0]['inline_data']['mime_type']}'"
)
def test_bedrock_tools_unpack_defs():
"""
Test that the unpack_defs method handles nested $ref inside anyOf items correctly
@ -2114,3 +2286,56 @@ def test_sanitize_messages_combined_case_a_and_case_d():
)
finally:
litellm.modify_params = original
def test_anthropic_messages_pt_file_block_preserves_cache_control():
"""
Test that cache_control is preserved on file-type content blocks
when translated to Anthropic document params.
Regression test for https://github.com/BerriAI/litellm/issues/23873
"""
from litellm.litellm_core_utils.prompt_templates.factory import (
anthropic_messages_pt,
)
messages = [
{
"role": "user",
"content": [
{
"type": "file",
"file": {
"filename": "doc.pdf",
"file_data": "data:application/pdf;base64,JVBERi0xLjQ=",
},
"cache_control": {"type": "ephemeral"},
},
{
"type": "text",
"text": "Summarize this document.",
"cache_control": {"type": "ephemeral"},
},
],
}
]
result = anthropic_messages_pt(
messages, model="claude-sonnet-4-20250514", llm_provider="anthropic"
)
content_blocks = result[0]["content"]
assert len(content_blocks) == 2
# Document block (from file) should preserve cache_control
doc_block = content_blocks[0]
assert doc_block["type"] == "document"
assert "cache_control" in doc_block, (
"cache_control was dropped from file/document block"
)
assert doc_block["cache_control"]["type"] == "ephemeral"
# Text block should also preserve cache_control
text_block = content_blocks[1]
assert text_block["type"] == "text"
assert "cache_control" in text_block
assert text_block["cache_control"]["type"] == "ephemeral"

View file

@ -73,6 +73,9 @@ class TestMapFinishReasonAnthropic:
def test_anthropic_finish_reasons(self, provider_reason: str, expected: str) -> None:
assert map_finish_reason(provider_reason) == expected
def test_refusal(self):
assert map_finish_reason("refusal") == "content_filter"
class TestMapFinishReasonGemini:
@pytest.mark.parametrize(

View file

@ -192,10 +192,11 @@ def test_azure_gpt5_1_series_temperature_handling(config: AzureOpenAIGPT5Config)
assert params["temperature"] == 0.6
def test_azure_gpt5_4_drops_reasoning_effort_when_tools_present(config: AzureOpenAIGPT5Config):
"""Azure Chat Completions: gpt-5.4+ drops reasoning_effort when tools are present.
def test_azure_gpt5_4_preserves_reasoning_effort_when_tools_present(config: AzureOpenAIGPT5Config):
"""Azure GPT-5.4+ no longer drops reasoning_effort when tools are present.
OpenAI routes tools+reasoning to Responses API; Azure does not, so we drop reasoning_effort.
Both OpenAI and Azure now route tools+reasoning to the Responses API bridge,
so reasoning_effort must be preserved in map_openai_params.
"""
tools = [{"type": "function", "function": {"name": "test", "description": "test"}}]
params = config.map_openai_params(
@ -205,7 +206,7 @@ def test_azure_gpt5_4_drops_reasoning_effort_when_tools_present(config: AzureOpe
drop_params=False,
api_version="2024-05-01-preview",
)
assert "reasoning_effort" not in params
assert params.get("reasoning_effort") == "high"
assert params["tools"] == tools

View file

@ -110,6 +110,35 @@ def test_get_supported_openai_params_reasoning_effort():
assert "reasoning_effort" not in unsupported_params
def test_add_transform_inline_image_block_skips_data_urls():
"""
data: URLs must not have #transform=inline appended — doing so corrupts the
base64 payload and raises binascii.Error: Incorrect padding on the Fireworks side.
Regression test for https://github.com/BerriAI/litellm/issues/23583
"""
config = FireworksAIConfig()
data_url = "data:image/jpeg;base64,/9j/4AAQSkZJRgAB"
# str branch
str_content = {"type": "image_url", "image_url": data_url}
result = config._add_transform_inline_image_block(
str_content, model="gpt-4", disable_add_transform_inline_image_block=False
)
assert result["image_url"] == data_url, "data URL must not be modified (str branch)"
# dict branch
dict_content = {"type": "image_url", "image_url": {"url": data_url}}
result = config._add_transform_inline_image_block(
dict_content, model="gpt-4", disable_add_transform_inline_image_block=False
)
assert result["image_url"]["url"] == data_url, "data URL must not be modified (dict branch)"
# regular https URL should still get the suffix
https_content = {"type": "image_url", "image_url": "https://example.com/image.jpg"}
result = config._add_transform_inline_image_block(
https_content, model="gpt-4", disable_add_transform_inline_image_block=False
)
assert result["image_url"].endswith("#transform=inline"), "https URL should get #transform=inline"
@pytest.mark.parametrize(
"api_base, expected_url_prefix",
[

View file

@ -158,6 +158,50 @@ def test_mistral_audio_transcription_response_transform():
assert response.text == "Four score and seven years ago..."
def test_mistral_audio_transcription_response_transform_diarized():
"""Test that diarized responses preserve segments and language."""
config = MistralAudioTranscriptionConfig()
mock_response = MagicMock(spec=httpx.Response)
mock_response.json.return_value = {
"model": "voxtral-mini-latest",
"text": "Hello, how are you? I am fine.",
"language": None,
"segments": [
{
"text": "Hello, how are you?",
"start": 0.3,
"end": 2.1,
"speaker_id": "speaker_1",
"type": "transcription_segment",
},
{
"text": "I am fine.",
"start": 2.5,
"end": 3.8,
"speaker_id": "speaker_2",
"type": "transcription_segment",
},
],
"usage": {
"prompt_audio_seconds": 4,
"prompt_tokens": 5,
"total_tokens": 50,
"completion_tokens": 20,
},
}
response = config.transform_audio_transcription_response(mock_response)
assert isinstance(response, TranscriptionResponse)
assert response.text == "Hello, how are you? I am fine."
assert response["segments"] is not None
assert len(response["segments"]) == 2
assert response["segments"][0]["speaker_id"] == "speaker_1"
assert response["segments"][1]["speaker_id"] == "speaker_2"
assert response["language"] is None
def test_mistral_audio_transcription_response_transform_empty():
config = MistralAudioTranscriptionConfig()

View file

@ -1317,4 +1317,41 @@ class TestVertexAIGlobalLocation:
# Assert correct URL format for global with beta API
expected_url = "https://aiplatform.googleapis.com/v1beta1/projects/test-project/locations/global/cachedContents"
assert url == expected_url, f"Expected {expected_url}, got {url}"
assert "global-aiplatform" not in url, "URL should not contain 'global-aiplatform' prefix"
assert "global-aiplatform" not in url, "URL should not contain 'global-aiplatform' prefix"
def test_gemini_context_caching_with_custom_api_base_passes_model(self):
"""Gemini context caching with custom api_base must pass model to _check_custom_proxy.
Regression test for https://github.com/BerriAI/litellm/issues/23846
Previously model was hardcoded to None, causing ValueError when api_base was set.
"""
caching = ContextCachingEndpoints()
auth_header, url = caching._get_token_and_url_context_caching(
gemini_api_key="test-key",
custom_llm_provider="gemini",
api_base="https://my-proxy.example.com",
vertex_project=None,
vertex_location=None,
vertex_auth_header=None,
model="gemini-1.5-pro",
)
assert "models/gemini-1.5-pro" in url
assert url.startswith("https://my-proxy.example.com/")
def test_gemini_context_caching_without_api_base_ignores_model(self):
"""Without custom api_base, model param is not needed (default URL is used)."""
caching = ContextCachingEndpoints()
auth_header, url = caching._get_token_and_url_context_caching(
gemini_api_key="test-key",
custom_llm_provider="gemini",
api_base=None,
vertex_project=None,
vertex_location=None,
vertex_auth_header=None,
)
assert "generativelanguage.googleapis.com" in url
assert "cachedContents" in url

View file

@ -230,3 +230,75 @@ def test_streaming_content_filter_finish_reason_preserved():
assert response is not None
assert len(response.choices) == 1
assert response.choices[0].finish_reason == "content_filter"
def test_streaming_tool_call_finish_reason_with_empty_content_in_final_chunk():
"""
When Gemini streams tool calls and the final chunk has BOTH empty content
(e.g. parts: [{text: ""}]) AND finishReason="STOP", the finish_reason
must still be "tool_calls".
This covers models like gemini-3.1-flash-lite-preview that send the
final chunk with content (empty text) instead of omitting it entirely.
Ref: https://github.com/BerriAI/litellm/issues/22900
"""
logging_obj = _make_logging_obj()
iterator = ModelResponseIterator(
streaming_response=iter([]),
sync_stream=True,
logging_obj=logging_obj,
)
# Chunk 1: tool call with no finishReason
chunk_with_tool_calls = {
"candidates": [
{
"content": {
"parts": [
{
"functionCall": {
"name": "get_weather",
"args": {"location": "San Francisco"},
}
}
],
"role": "model",
},
"index": 0,
}
],
}
# Chunk 2: finishReason="STOP" WITH empty content (text: "")
chunk_with_empty_content_and_finish = {
"candidates": [
{
"content": {
"parts": [{"text": ""}],
"role": "model",
},
"finishReason": "STOP",
"index": 0,
}
],
"usageMetadata": {
"promptTokenCount": 50,
"candidatesTokenCount": 20,
"totalTokenCount": 70,
},
}
# Process chunk 1
response1 = iterator.chunk_parser(chunk_with_tool_calls)
assert response1 is not None
assert len(response1.choices) == 1
assert response1.choices[0].delta.tool_calls is not None
assert iterator.has_seen_tool_calls is True
# Process chunk 2 (final chunk with empty content)
response2 = iterator.chunk_parser(chunk_with_empty_content_and_finish)
assert response2 is not None
assert len(response2.choices) == 1
# Must be "tool_calls", NOT "stop"
assert response2.choices[0].finish_reason == "tool_calls"

View file

@ -0,0 +1,164 @@
"""
Tests for Vertex AI partner models count_tokens location resolution.
Ref: https://github.com/BerriAI/litellm/issues/23872
"""
import pytest
from litellm.llms.vertex_ai.vertex_ai_partner_models.count_tokens.handler import (
VertexAIPartnerModelsTokenCounter,
)
@pytest.fixture
def counter():
return VertexAIPartnerModelsTokenCounter()
class TestCountTokensLocationResolution:
"""Verify that vertex_count_tokens_location is respected in handle_count_tokens_request."""
def _build_litellm_params(
self,
vertex_location=None,
vertex_count_tokens_location=None,
):
params = {}
if vertex_location is not None:
params["vertex_location"] = vertex_location
if vertex_count_tokens_location is not None:
params["vertex_count_tokens_location"] = vertex_count_tokens_location
return params
@pytest.mark.asyncio
async def test_count_tokens_location_overrides_vertex_location(self, counter, monkeypatch):
"""vertex_count_tokens_location should take precedence over vertex_location."""
captured = {}
async def fake_ensure_access_token(self, credentials, project_id, custom_llm_provider):
return "fake-token", "fake-project"
def fake_build_endpoint(self, model, project_id, vertex_location, api_base=None):
captured["vertex_location"] = vertex_location
return "https://fake-endpoint"
monkeypatch.setattr(
VertexAIPartnerModelsTokenCounter, "_ensure_access_token_async", fake_ensure_access_token
)
monkeypatch.setattr(
VertexAIPartnerModelsTokenCounter, "_build_count_tokens_endpoint", fake_build_endpoint
)
# Mock the HTTP call to avoid real network requests
class FakeResponse:
status_code = 200
def json(self):
return {"input_tokens": 10}
def raise_for_status(self):
pass
class FakeClient:
async def post(self, url, headers=None, json=None, **kwargs):
return FakeResponse()
import litellm.llms.vertex_ai.vertex_ai_partner_models.count_tokens.handler as handler_mod
monkeypatch.setattr(handler_mod, "get_async_httpx_client", lambda **kwargs: FakeClient())
litellm_params = self._build_litellm_params(
vertex_location="us-east5",
vertex_count_tokens_location="europe-west1",
)
await counter.handle_count_tokens_request(
model="claude-sonnet-4-6",
request_data={"messages": [{"role": "user", "content": "hi"}]},
litellm_params=litellm_params,
)
assert captured["vertex_location"] == "europe-west1"
@pytest.mark.asyncio
async def test_claude_without_count_tokens_location_defaults_to_us_east5(self, counter, monkeypatch):
"""Claude models without any location should default to us-east5."""
captured = {}
async def fake_ensure_access_token(self, credentials, project_id, custom_llm_provider):
return "fake-token", "fake-project"
def fake_build_endpoint(self, model, project_id, vertex_location, api_base=None):
captured["vertex_location"] = vertex_location
return "https://fake-endpoint"
monkeypatch.setattr(
VertexAIPartnerModelsTokenCounter, "_ensure_access_token_async", fake_ensure_access_token
)
monkeypatch.setattr(
VertexAIPartnerModelsTokenCounter, "_build_count_tokens_endpoint", fake_build_endpoint
)
class FakeResponse:
status_code = 200
def json(self):
return {"input_tokens": 10}
def raise_for_status(self):
pass
class FakeClient:
async def post(self, url, headers=None, json=None, **kwargs):
return FakeResponse()
import litellm.llms.vertex_ai.vertex_ai_partner_models.count_tokens.handler as handler_mod
monkeypatch.setattr(handler_mod, "get_async_httpx_client", lambda **kwargs: FakeClient())
litellm_params = self._build_litellm_params() # no location at all
await counter.handle_count_tokens_request(
model="claude-sonnet-4-6",
request_data={"messages": [{"role": "user", "content": "hi"}]},
litellm_params=litellm_params,
)
assert captured["vertex_location"] == "us-east5"
@pytest.mark.asyncio
async def test_claude_with_vertex_location_uses_it(self, counter, monkeypatch):
"""Claude models with vertex_location but no count_tokens_location should use vertex_location."""
captured = {}
async def fake_ensure_access_token(self, credentials, project_id, custom_llm_provider):
return "fake-token", "fake-project"
def fake_build_endpoint(self, model, project_id, vertex_location, api_base=None):
captured["vertex_location"] = vertex_location
return "https://fake-endpoint"
monkeypatch.setattr(
VertexAIPartnerModelsTokenCounter, "_ensure_access_token_async", fake_ensure_access_token
)
monkeypatch.setattr(
VertexAIPartnerModelsTokenCounter, "_build_count_tokens_endpoint", fake_build_endpoint
)
class FakeResponse:
status_code = 200
def json(self):
return {"input_tokens": 10}
def raise_for_status(self):
pass
class FakeClient:
async def post(self, url, headers=None, json=None, **kwargs):
return FakeResponse()
import litellm.llms.vertex_ai.vertex_ai_partner_models.count_tokens.handler as handler_mod
monkeypatch.setattr(handler_mod, "get_async_httpx_client", lambda **kwargs: FakeClient())
litellm_params = self._build_litellm_params(vertex_location="asia-southeast1")
await counter.handle_count_tokens_request(
model="claude-sonnet-4-6",
request_data={"messages": [{"role": "user", "content": "hi"}]},
litellm_params=litellm_params,
)
assert captured["vertex_location"] == "asia-southeast1"

View file

@ -0,0 +1,49 @@
"""
Test that max_budget from environment variable (string) is correctly
converted to float.
GitHub Issue: #23843
"""
from unittest.mock import patch
import pytest
import litellm
@pytest.mark.asyncio
async def test_max_budget_string_converted_to_float():
"""
When max_budget is set via os.environ/MAX_BUDGET, it arrives as a
string. initialize() should convert it to float so the comparison
`litellm.max_budget > 0` doesn't raise TypeError.
"""
with patch("litellm.proxy.common_utils.banner.show_banner"), patch(
"litellm.proxy.proxy_server.generate_feedback_box"
):
from litellm.proxy.proxy_server import initialize
original = litellm.max_budget
try:
await initialize(max_budget="100.5")
assert isinstance(litellm.max_budget, float)
assert litellm.max_budget == 100.5
finally:
litellm.max_budget = original
@pytest.mark.asyncio
async def test_max_budget_float_stays_float():
"""max_budget as float should still work."""
with patch("litellm.proxy.common_utils.banner.show_banner"), patch(
"litellm.proxy.proxy_server.generate_feedback_box"
):
from litellm.proxy.proxy_server import initialize
original = litellm.max_budget
try:
await initialize(max_budget=200.0)
assert isinstance(litellm.max_budget, float)
assert litellm.max_budget == 200.0
finally:
litellm.max_budget = original

View file

@ -661,6 +661,40 @@ def test_responses_api_bridge_check_gpt_5_5_tools_plus_reasoning_routes_to_respo
assert model_info.get("mode") == "responses"
def test_responses_api_bridge_check_azure_gpt_5_4_tools_plus_reasoning_routes_to_responses():
"""Azure gpt-5.4 with both tools and reasoning_effort should route to Responses API."""
from litellm.main import responses_api_bridge_check
with patch("litellm.main._get_model_info_helper") as mock_get_model_info:
mock_get_model_info.return_value = {"max_tokens": 128000}
model_info, model = responses_api_bridge_check(
model="gpt-5.4",
custom_llm_provider="azure",
tools=[{"type": "function", "function": {"name": "get_capital"}}],
reasoning_effort="high",
)
assert model == "gpt-5.4"
assert model_info.get("mode") == "responses"
def test_responses_api_bridge_check_azure_gpt_5_4_tools_without_reasoning_stays_chat():
"""Azure gpt-5.4 with tools only should not be force-routed to Responses API."""
from litellm.main import responses_api_bridge_check
with patch("litellm.main._get_model_info_helper") as mock_get_model_info:
mock_get_model_info.return_value = {"max_tokens": 128000}
model_info, model = responses_api_bridge_check(
model="gpt-5.4",
custom_llm_provider="azure",
tools=[{"type": "function", "function": {"name": "get_capital"}}],
reasoning_effort=None,
)
assert model == "gpt-5.4"
assert model_info.get("mode") != "responses"
def test_responses_api_bridge_check_gpt_5_4_tools_without_reasoning_stays_chat():
"""gpt-5.4 with tools only should not be force-routed to Responses API."""
from litellm.main import responses_api_bridge_check

View file

@ -1,7 +1,15 @@
from litellm._redis import get_redis_url_from_environment, _get_redis_cluster_kwargs, get_redis_async_client
from litellm._redis import (
get_redis_url_from_environment,
_get_redis_cluster_kwargs,
get_redis_async_client,
get_redis_client,
get_redis_connection_pool,
)
import json
import os
import pytest
from unittest.mock import MagicMock, patch
import redis
import redis.asyncio as async_redis
def test_get_redis_url_from_environment_single_url(monkeypatch):
@ -167,3 +175,115 @@ def test_get_redis_async_client_without_connection_pool():
# Verify Redis was called without connection_pool in kwargs
call_kwargs = mock_redis.call_args[1]
assert "connection_pool" not in call_kwargs, "connection_pool should not be in kwargs when not provided"
@patch("litellm._redis.init_redis_cluster")
def test_sync_client_prefers_cluster_over_url(mock_init_cluster, monkeypatch):
"""
Test get_redis_client returns RedisCluster when startup_nodes is present even if
REDIS_URL is also set.
"""
monkeypatch.setenv("REDIS_URL", "redis://fallback-host:6379")
mock_init_cluster.return_value = MagicMock(spec=redis.RedisCluster)
startup_nodes = [{"host": "cluster-node.example.com", "port": 6379}]
get_redis_client(startup_nodes=startup_nodes)
mock_init_cluster.assert_called_once()
call_kwargs = mock_init_cluster.call_args[0][0]
assert (
"startup_nodes" in call_kwargs
), "startup_nodes must be forwarded to init_redis_cluster"
@patch("litellm._redis.async_redis.RedisCluster")
def test_async_client_prefers_cluster_over_url(mock_cluster_cls, monkeypatch):
"""
Test (1) get_redis_async_client returns async RedisCluster when startup_nodes is present
even if REDIS_URL is also set and (2) startup_nodes is forwarded to RedisCluster.
"""
monkeypatch.setenv("REDIS_URL", "redis://fallback-host:6379")
startup_nodes = [{"host": "cluster-node.example.com", "port": 6379}]
get_redis_async_client(startup_nodes=startup_nodes)
mock_cluster_cls.assert_called_once()
call_kwargs = mock_cluster_cls.call_args[1]
assert "startup_nodes" in call_kwargs, "startup_nodes must be forwarded to async RedisCluster"
assert len(call_kwargs["startup_nodes"]) == 1, "should forward exactly 1 cluster node"
@patch("litellm._redis.async_redis.RedisCluster")
def test_async_client_prefers_cluster_over_url_via_env_var(mock_cluster_cls, monkeypatch):
"""
Test get_redis_async_client returns async RedisCluster when REDIS_CLUSTER_NODES is set
even if REDIS_URL is also set.
"""
monkeypatch.setenv("REDIS_URL", "redis://fallback-host:6379")
monkeypatch.setenv(
"REDIS_CLUSTER_NODES",
json.dumps([{"host": "cluster-node.example.com", "port": 6379}]),
)
get_redis_async_client()
mock_cluster_cls.assert_called_once()
call_kwargs = mock_cluster_cls.call_args[1]
assert "startup_nodes" in call_kwargs, "startup_nodes must be forwarded to async RedisCluster"
@patch("litellm._redis.init_redis_cluster")
def test_sync_client_prefers_cluster_over_url_via_env_var(mock_init_cluster, monkeypatch):
"""
Test get_redis_client returns RedisCluster when REDIS_CLUSTER_NODES is set even if
REDIS_URL is also set.
"""
monkeypatch.setenv("REDIS_URL", "redis://fallback-host:6379")
monkeypatch.setenv(
"REDIS_CLUSTER_NODES",
json.dumps([{"host": "cluster-node.example.com", "port": 6379}]),
)
mock_init_cluster.return_value = MagicMock(spec=redis.RedisCluster)
get_redis_client()
mock_init_cluster.assert_called_once()
call_kwargs = mock_init_cluster.call_args[0][0]
assert "startup_nodes" in call_kwargs, "startup_nodes must be forwarded to init_redis_cluster"
assert len(call_kwargs["startup_nodes"]) == 1
@patch("litellm._redis.init_redis_cluster")
def test_sync_client_preserves_password_for_cluster_when_url_also_set(mock_init_cluster, monkeypatch):
"""
Test _get_redis_client_logic does not strip password from redis_kwargs when
startup_nodes is present even if REDIS_URL is also set.
"""
monkeypatch.setenv("REDIS_URL", "redis://fallback-host:6379")
monkeypatch.setenv("REDIS_PASSWORD", "secret")
mock_init_cluster.return_value = MagicMock(spec=redis.RedisCluster)
startup_nodes = [{"host": "cluster-node.example.com", "port": 6379}]
get_redis_client(startup_nodes=startup_nodes)
mock_init_cluster.assert_called_once()
call_kwargs = mock_init_cluster.call_args[0][0]
assert "password" in call_kwargs, "password must not be stripped when routing to cluster"
assert call_kwargs["password"] == "secret"
def test_connection_pool_returns_none_for_cluster(monkeypatch):
"""Test get_redis_connection_pool returns None when startup_nodes is present."""
monkeypatch.setenv("REDIS_URL", "redis://fallback-host:6379")
startup_nodes = [{"host": "cluster-node.example.com", "port": 6379}]
result = get_redis_connection_pool(startup_nodes=startup_nodes)
assert result is None, "connection pool must be None for cluster mode"
@patch("litellm._redis.redis.Redis.from_url")
def test_sync_client_url_used_when_no_cluster(mock_from_url, monkeypatch):
"""
Test get_redis_client default to using URL path when no startup_nodes are provided.
"""
monkeypatch.setenv("REDIS_URL", "redis://plain-host:6379")
monkeypatch.delenv("REDIS_CLUSTER_NODES", raising=False)
get_redis_client()
mock_from_url.assert_called_once()

View file

@ -0,0 +1,10 @@
<svg width="20" height="20" viewBox="0 0 20 20" fill="none" xmlns="http://www.w3.org/2000/svg">
<path d="M18.0216 0C17.8264 0 17.5923 0.078125 17.3581 0.15625L1.12266 7.65625C-0.750664 8.51562 -0.126225 11.2891 1.86418 11.2891H8.73302V18.1641C8.73302 19.3359 9.7087 20 10.6454 20C11.3088 20 12.0113 19.6875 12.3626 18.9062L19.8559 2.65625C20.4022 1.36719 19.3095 0 18.0216 0ZM18.7241 2.14844L14.6743 10.625L11.2308 18.3594C11.1137 18.6328 10.9186 18.7891 10.6454 18.7891C10.4112 18.7891 9.9819 18.6328 9.9819 18.1641V10H1.87332C1.5611 10 1.24888 9.6875 1.24888 9.375C1.24888 9.0625 1.5611 8.75 1.87332 8.75L14.6743 10.625L17.8264 1.32812C17.9045 1.28906 17.9825 1.25 18.0216 1.25C18.2557 1.25 18.4899 1.40625 18.646 1.64062C18.7241 1.75781 18.8021 1.95312 18.7241 2.14844Z" fill="url(#paint0_linear_22912_90694)"/>
<defs>
<linearGradient id="paint0_linear_22912_90694" x1="20" y1="9.56158e-07" x2="3.95833" y2="16.25" gradientUnits="userSpaceOnUse">
<stop stop-color="#D500F9"/>
<stop offset="0.5" stop-color="#6200EA"/>
<stop offset="1" stop-color="#2E006D"/>
</linearGradient>
</defs>
</svg>

After

Width:  |  Height:  |  Size: 1.1 KiB

View file

@ -264,4 +264,10 @@ export const GUARDRAIL_PRESETS: Record<string, GuardrailPreset> = {
mode: "pre_call",
defaultOn: false,
},
akto: {
provider: "Akto",
guardrailNameSuggestion: "Akto Guardrail",
mode: "pre_call",
defaultOn: false,
},
};

View file

@ -373,6 +373,14 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [
logo: `${ASSET_PREFIX}pillar.jpeg`,
tags: ["Monitoring", "Safety"],
},
{
id: "akto",
name: "Akto Guardrail",
description: "AI security platform from Akto.io with automatic monitoring and guardrails for AI/ML applications.",
category: "partner",
logo: `${ASSET_PREFIX}akto.svg`,
tags: ["Security", "Safety", "Monitoring"],
},
];
export const ALL_CARDS = [...LITELLM_CONTENT_FILTER_CARDS, ...PARTNER_GUARDRAIL_CARDS];

View file

@ -125,6 +125,7 @@ export const guardrailLogoMap: Record<string, string> = {
EnkryptAI: `${asset_logos_folder}enkrypt_ai.avif`,
"Prompt Security": `${asset_logos_folder}prompt_security.png`,
"LiteLLM Content Filter": `${asset_logos_folder}litellm_logo.jpg`,
"Akto": `${asset_logos_folder}akto.svg`,
};
export const getGuardrailLogoAndName = (guardrailValue: string): { logo: string; displayName: string } => {