mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
commit
f911d8d865
41 changed files with 2768 additions and 265 deletions
139
docs/my-website/docs/proxy/guardrails/akto.md
Normal file
139
docs/my-website/docs/proxy/guardrails/akto.md
Normal 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) |
|
||||
|
|
@ -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
|
||||
```
|
||||
|
||||
:::
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
37
litellm/proxy/guardrails/guardrail_hooks/akto/__init__.py
Normal file
37
litellm/proxy/guardrails/guardrail_hooks/akto/__init__.py
Normal 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,
|
||||
}
|
||||
456
litellm/proxy/guardrails/guardrail_hooks/akto/akto.py
Normal file
456
litellm/proxy/guardrails/guardrail_hooks/akto/akto.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
55
litellm/types/proxy/guardrails/guardrail_hooks/akto.py
Normal file
55
litellm/types/proxy/guardrails/guardrail_hooks/akto.py
Normal 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"
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
550
tests/guardrails_tests/test_akto_guardrails.py
Normal file
550
tests/guardrails_tests/test_akto_guardrails.py
Normal 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"
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
49
tests/test_litellm/proxy/test_max_budget_env_var.py
Normal file
49
tests/test_litellm/proxy/test_max_budget_env_var.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
10
ui/litellm-dashboard/public/assets/logos/akto.svg
Normal file
10
ui/litellm-dashboard/public/assets/logos/akto.svg
Normal 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 |
|
|
@ -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,
|
||||
},
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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];
|
||||
|
|
|
|||
|
|
@ -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 } => {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue