Merge branch 'main' into doc_yj_feb21

This commit is contained in:
yuneng-jiang
2026-02-21 16:00:46 -08:00
committed by GitHub
32 changed files with 2039 additions and 347 deletions
@@ -574,6 +574,8 @@ router_settings:
| DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO | Default minimal reasoning effort thinking budget for Gemini 2.5 Pro. Default is 512
| DEFAULT_REDIS_MAJOR_VERSION | Default Redis major version to assume when version cannot be determined. Default is 7
| DEFAULT_REDIS_SYNC_INTERVAL | Default Redis synchronization interval in seconds. Default is 1
| DEFAULT_SEMANTIC_GUARD_EMBEDDING_MODEL | Default embedding model for Semantic Guard (route-matching guardrail). Default is "text-embedding-3-small"
| DEFAULT_SEMANTIC_GUARD_SIMILARITY_THRESHOLD | Default similarity threshold for Semantic Guard route matching. Default is 0.75
| DEFAULT_REPLICATE_GPU_PRICE_PER_SECOND | Default price per second for Replicate GPU. Default is 0.001400
| DEFAULT_REPLICATE_POLLING_DELAY_SECONDS | Default delay in seconds for Replicate polling. Default is 1
| DEFAULT_REPLICATE_POLLING_RETRIES | Default number of retries for Replicate polling. Default is 5
@@ -756,6 +758,7 @@ router_settings:
| LITELLM_CLI_JWT_EXPIRATION_HOURS | Expiration time in hours for CLI-generated JWT tokens. Default is 24 hours
| LITELLM_DD_AGENT_HOST | Hostname or IP of DataDog agent for LiteLLM-specific logging. When set, logs are sent to agent instead of direct API
| LITELLM_DEPLOYMENT_ENVIRONMENT | Environment name for the deployment (e.g., "production", "staging"). Used as a fallback when OTEL_ENVIRONMENT_NAME is not set. Sets the `environment` tag in telemetry data
| LITELLM_DETAILED_TIMING | When true, adds detailed per-phase timing headers to responses (x-litellm-timing-{pre-processing,llm-api,post-processing,message-copy}-ms). Default is false. See [latency overhead docs](../troubleshoot/latency_overhead.md)
| LITELLM_DD_AGENT_PORT | Port of DataDog agent for LiteLLM-specific log intake. Default is 10518
| LITELLM_DD_LLM_OBS_PORT | Port for Datadog LLM Observability agent. Default is 8126
| LITELLM_DONT_SHOW_FEEDBACK_BOX | Flag to hide feedback box in LiteLLM UI
@@ -807,6 +810,7 @@ router_settings:
| LOGGING_WORKER_MAX_QUEUE_SIZE | Maximum size of the logging worker queue. When the queue is full, the worker aggressively clears tasks to make room instead of dropping logs. Default is 50,000
| LOGGING_WORKER_MAX_TIME_PER_COROUTINE | Maximum time in seconds allowed for each coroutine in the logging worker before timing out. Default is 20.0
| LOGGING_WORKER_CLEAR_PERCENTAGE | Percentage of the queue to extract when clearing. Default is 50%
| MAX_BASE64_LENGTH_FOR_LOGGING | Maximum number of base64 characters to keep in logging payloads. Data URIs exceeding this are replaced with a size placeholder. Set to 0 to disable truncation. Default is 64
| MAX_COMPETITOR_NAMES | Maximum number of competitor names allowed in policy template enrichment. Default is 100
| MAX_EXCEPTION_MESSAGE_LENGTH | Maximum length for exception messages. Default is 2000
| MAX_ITERATIONS_TO_CLEAR_QUEUE | Maximum number of iterations to attempt when clearing the logging worker queue during shutdown. Default is 200
@@ -830,6 +834,7 @@ router_settings:
| MAX_LANGFUSE_INITIALIZED_CLIENTS | Maximum number of Langfuse clients to initialize on proxy. Default is 50. This is set since langfuse initializes 1 thread everytime a client is initialized. We've had an incident in the past where we reached 100% cpu utilization because Langfuse was initialized several times.
| MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH | Maximum header length for MCP semantic filter tools. Default is 150
| MAX_POLICY_ESTIMATE_IMPACT_ROWS | Maximum number of rows returned when estimating the impact of a policy. Default is 1000
| MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG | Maximum payload size in bytes for full DEBUG serialization. Payloads exceeding this will be truncated in logs. Default is 102400 (100 KB)
| MIN_NON_ZERO_TEMPERATURE | Minimum non-zero temperature value. Default is 0.0001
| MINIMUM_PROMPT_CACHE_TOKEN_COUNT | Minimum token count for caching a prompt. Default is 1024
| MISTRAL_API_BASE | Base URL for Mistral API. Default is https://api.mistral.ai
@@ -893,6 +898,13 @@ router_settings:
| POSTHOG_API_URL | Base URL for PostHog API (defaults to https://us.i.posthog.com)
| POSTHOG_MOCK | Enable mock mode for PostHog integration testing. When set to true, intercepts PostHog API calls and returns mock responses without making actual network calls. Default is false
| POSTHOG_MOCK_LATENCY_MS | Mock latency in milliseconds for PostHog API calls when mock mode is enabled. Simulates network round-trip time. Default is 100ms
| PRISMA_AUTH_RECONNECT_LOCK_TIMEOUT_SECONDS | Lock timeout in seconds for Prisma auth reconnection. Default is 0.1
| PRISMA_AUTH_RECONNECT_TIMEOUT_SECONDS | Timeout in seconds for Prisma auth reconnection attempts. Default is 2.0
| PRISMA_HEALTH_WATCHDOG_ENABLED | Enable the Prisma DB health watchdog that monitors and reconnects on connection loss. Default is true
| PRISMA_HEALTH_WATCHDOG_INTERVAL_SECONDS | Interval in seconds for Prisma health watchdog probes. Default is 30
| PRISMA_HEALTH_WATCHDOG_PROBE_TIMEOUT_SECONDS | Timeout in seconds for each Prisma health probe. Default is 5.0
| PRISMA_RECONNECT_COOLDOWN_SECONDS | Cooldown in seconds between Prisma reconnection attempts. Default is 15
| PRISMA_WATCHDOG_RECONNECT_TIMEOUT_SECONDS | Timeout in seconds for Prisma watchdog-initiated reconnection. Default is 30.0
| PREDIBASE_API_BASE | Base URL for Predibase API
| PRESIDIO_ANALYZER_API_BASE | Base URL for Presidio Analyzer service
| PRESIDIO_ANONYMIZER_API_BASE | Base URL for Presidio Anonymizer service
+9
View File
@@ -0,0 +1,9 @@
{
"tornado:6.5.3": "Apache-2.0",
"redisvl:0.4.1": "MIT",
"google-cloud-iam:2.19.1": "Apache 2.0",
"google-genai:1.37.0": "Apache-2.0",
"azure-keyvault:4.2.0": "MIT License",
"soundfile:0.12.1": "BSD 3-Clause License",
"openapi-core:0.21.0": "BSD-3-Clause"
}
+1
View File
@@ -98,6 +98,7 @@ _custom_logger_compatible_callbacks_literal = Literal[
"openmeter",
"logfire",
"literalai",
"litellm_agent",
"dynamic_rate_limiter",
"dynamic_rate_limiter_v3",
"langsmith",
@@ -0,0 +1,5 @@
"""LiteLLM Agent integration - model name resolver for litellm_agent/ prefix."""
from .litellm_agent_model_resolver import LiteLLMAgentModelResolver
__all__ = ["LiteLLMAgentModelResolver"]
@@ -0,0 +1,79 @@
"""
Hook for LiteLLM that strips the litellm_agent/ prefix from model names.
When model is litellm_agent/gpt-3.5-turbo, this hook replaces it with gpt-3.5-turbo
before the completion call, similar to langfuse/model resolution.
"""
from typing import Dict, List, Optional, Tuple
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.llms.openai import AllMessageValues
from litellm.types.prompts.init_prompts import PromptSpec
from litellm.types.utils import StandardCallbackDynamicParams
LITELLM_AGENT_PREFIX = "litellm_agent/"
class LiteLLMAgentModelResolver(CustomLogger):
"""
CustomLogger that strips litellm_agent/ prefix from model names.
Enables model configs like litellm_agent/gpt-3.5-turbo to resolve to gpt-3.5-turbo.
"""
def get_chat_completion_prompt(
self,
model: str,
messages: List[AllMessageValues],
non_default_params: dict,
prompt_id: Optional[str],
prompt_variables: Optional[dict],
dynamic_callback_params: StandardCallbackDynamicParams,
prompt_spec: Optional[PromptSpec] = None,
prompt_label: Optional[str] = None,
prompt_version: Optional[int] = None,
ignore_prompt_manager_model: Optional[bool] = False,
ignore_prompt_manager_optional_params: Optional[bool] = False,
) -> Tuple[str, List[AllMessageValues], dict]:
"""
Strip litellm_agent/ prefix from model name.
Returns:
(resolved_model, messages, non_default_params)
"""
if ignore_prompt_manager_model:
return model, messages, non_default_params
resolved_model = model.replace(LITELLM_AGENT_PREFIX, "", 1)
return resolved_model, messages, non_default_params
async def async_get_chat_completion_prompt(
self,
model: str,
messages: List[AllMessageValues],
non_default_params: dict,
prompt_id: Optional[str],
prompt_variables: Optional[dict],
dynamic_callback_params: StandardCallbackDynamicParams,
litellm_logging_obj: object,
prompt_spec: Optional[PromptSpec] = None,
tools: Optional[List[Dict]] = None,
prompt_label: Optional[str] = None,
prompt_version: Optional[int] = None,
ignore_prompt_manager_model: Optional[bool] = False,
ignore_prompt_manager_optional_params: Optional[bool] = False,
) -> Tuple[str, List[AllMessageValues], dict]:
"""Async delegate to get_chat_completion_prompt."""
return self.get_chat_completion_prompt(
model=model,
messages=messages,
non_default_params=non_default_params,
prompt_id=prompt_id,
prompt_variables=prompt_variables,
dynamic_callback_params=dynamic_callback_params,
prompt_spec=prompt_spec,
prompt_label=prompt_label,
prompt_version=prompt_version,
ignore_prompt_manager_model=ignore_prompt_manager_model,
ignore_prompt_manager_optional_params=ignore_prompt_manager_optional_params,
)
@@ -18,11 +18,11 @@ from litellm.integrations.azure_storage.azure_storage import AzureBlobStorageLog
from litellm.integrations.bitbucket import BitBucketPromptManager
from litellm.integrations.braintrust_logging import BraintrustLogger
from litellm.integrations.cloudzero.cloudzero import CloudZeroLogger
from litellm.integrations.focus.focus_logger import FocusLogger
from litellm.integrations.datadog.datadog import DataDogLogger
from litellm.integrations.datadog.datadog_llm_obs import DataDogLLMObsLogger
from litellm.integrations.deepeval import DeepEvalLogger
from litellm.integrations.dotprompt import DotpromptManager
from litellm.integrations.focus.focus_logger import FocusLogger
from litellm.integrations.galileo import GalileoObserve
from litellm.integrations.gcs_bucket.gcs_bucket import GCSBucketLogger
from litellm.integrations.gcs_pubsub.pub_sub import GcsPubSubLogger
@@ -33,6 +33,7 @@ from litellm.integrations.langfuse.langfuse_prompt_management import (
LangfusePromptManagement,
)
from litellm.integrations.langsmith import LangsmithLogger
from litellm.integrations.litellm_agent import LiteLLMAgentModelResolver
from litellm.integrations.literal_ai import LiteralAILogger
from litellm.integrations.mlflow import MlflowLogger
from litellm.integrations.openmeter import OpenMeterLogger
@@ -61,6 +62,7 @@ class CustomLoggerRegistry:
"galileo": GalileoObserve,
"langsmith": LangsmithLogger,
"literalai": LiteralAILogger,
"litellm_agent": LiteLLMAgentModelResolver,
"prometheus": PrometheusLogger,
"datadog": DataDogLogger,
"datadog_llm_observability": DataDogLLMObsLogger,
@@ -147,6 +147,7 @@ from ..integrations.langfuse.langfuse import LangFuseLogger
from ..integrations.langfuse.langfuse_handler import LangFuseHandler
from ..integrations.langfuse.langfuse_prompt_management import LangfusePromptManagement
from ..integrations.langsmith import LangsmithLogger
from ..integrations.litellm_agent import LiteLLMAgentModelResolver
from ..integrations.literal_ai import LiteralAILogger
from ..integrations.logfire_logger import LogfireLevel, LogfireLogger
from ..integrations.lunary import LunaryLogger
@@ -587,6 +588,11 @@ class Logging(LiteLLMLoggingBaseClass):
if prompt_id:
return True
# Check if model uses litellm_agent prefix (model replacement without prompt_id)
model = non_default_params.get("model", "")
if isinstance(model, str) and model.startswith("litellm_agent/"):
return True
if self._should_run_prompt_management_hooks_without_prompt_id(
non_default_params=non_default_params,
tools=tools,
@@ -3629,6 +3635,14 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
_literalai_logger = LiteralAILogger()
_in_memory_loggers.append(_literalai_logger)
return _literalai_logger # type: ignore
elif logging_integration == "litellm_agent":
for callback in _in_memory_loggers:
if isinstance(callback, LiteLLMAgentModelResolver):
return callback # type: ignore
_litellm_agent_resolver = LiteLLMAgentModelResolver()
_in_memory_loggers.append(_litellm_agent_resolver)
return _litellm_agent_resolver # type: ignore
elif logging_integration == "prometheus":
PrometheusLogger = _get_cached_prometheus_logger()
@@ -4183,6 +4197,10 @@ def get_custom_logger_compatible_class( # noqa: PLR0915
for callback in _in_memory_loggers:
if isinstance(callback, LiteralAILogger):
return callback
elif logging_integration == "litellm_agent":
for callback in _in_memory_loggers:
if isinstance(callback, LiteLLMAgentModelResolver):
return callback
elif logging_integration == "prometheus":
PrometheusLogger = _get_cached_prometheus_logger()
for callback in _in_memory_loggers:
@@ -452,7 +452,7 @@ def update_responses_input_with_model_file_ids(
For managed files (unified file IDs), uses model_file_id_mapping if provided,
otherwise decodes the base64-encoded unified file ID and extracts the llm_output_file_id directly.
Args:
input: The responses API input parameter
model_id: The model ID to use for looking up provider-specific file IDs
@@ -488,9 +488,13 @@ def update_responses_input_with_model_file_ids(
file_id = content_item.get("file_id")
if file_id:
provider_file_id = file_id # Default to original
# Check if we have a mapping for this file ID
if model_file_id_mapping and model_id and file_id in model_file_id_mapping:
if (
model_file_id_mapping
and model_id
and file_id in model_file_id_mapping
):
# Use the model-specific file ID from mapping
provider_file_id = (
model_file_id_mapping.get(file_id, {}).get(model_id)
@@ -501,15 +505,19 @@ def update_responses_input_with_model_file_ids(
updated_content.append(updated_content_item)
else:
# Check if this is a base64-encoded unified file ID without mapping
is_unified_file_id = _is_base64_encoded_unified_file_id(file_id)
is_unified_file_id = _is_base64_encoded_unified_file_id(
file_id
)
if is_unified_file_id:
# Fallback: decode unified file ID
unified_file_id = convert_b64_uid_to_unified_uid(file_id)
unified_file_id = convert_b64_uid_to_unified_uid(
file_id
)
if "llm_output_file_id," in unified_file_id:
provider_file_id = unified_file_id.split(
"llm_output_file_id,"
)[1].split(";")[0]
updated_content_item = content_item.copy()
updated_content_item["file_id"] = provider_file_id
updated_content.append(updated_content_item)
@@ -534,9 +542,9 @@ def update_responses_tools_with_model_file_ids(
) -> Optional[List[Dict[str, Any]]]:
"""
Updates responses API tools with provider-specific file IDs.
Handles code_interpreter tools with container.file_ids.
Args:
tools: The responses API tools parameter
model_id: The model ID to use for looking up provider-specific file IDs
@@ -545,18 +553,18 @@ def update_responses_tools_with_model_file_ids(
"""
if not tools or not isinstance(tools, list):
return tools
if not model_file_id_mapping or not model_id:
return tools
updated_tools = []
for tool in tools:
if not isinstance(tool, dict):
updated_tools.append(tool)
continue
updated_tool = tool.copy()
# Handle code_interpreter with container file_ids
if tool.get("type") == "code_interpreter":
container = tool.get("container")
@@ -578,14 +586,14 @@ def update_responses_tools_with_model_file_ids(
updated_file_ids.append(file_id)
else:
updated_file_ids.append(file_id)
# Update the tool with new file IDs
updated_container = container.copy()
updated_container["file_ids"] = updated_file_ids
updated_tool["container"] = updated_container
updated_tools.append(updated_tool)
return updated_tools
@@ -1104,6 +1112,46 @@ def set_last_user_message(
return messages
def add_system_prompt_to_messages(
messages: List[AllMessageValues],
system_prompt: str,
merge_with_first_system: bool = False,
) -> List[AllMessageValues]:
"""
Add a system prompt to the messages list.
Args:
messages: List of chat completion messages
system_prompt: The system prompt content to add. If empty or None, returns messages unchanged.
merge_with_first_system: If True and the first message is already a system message,
prepends the new prompt to that message's content. If False, adds a new system
message at the beginning.
Returns:
New list of messages with the system prompt added
"""
if not system_prompt:
return list(messages)
if merge_with_first_system and messages and messages[0].get("role") == "system":
first = dict(messages[0])
existing_content = first.get("content", "")
merged_content: Union[str, List[Dict[str, str]]]
if isinstance(existing_content, str):
merged_content = f"{system_prompt.strip()}\n\n{existing_content}"
elif isinstance(existing_content, list):
merged_content = [{"type": "text", "text": system_prompt.strip()}] + list(
existing_content
)
else:
merged_content = [{"type": "text", "text": system_prompt.strip()}]
first["content"] = merged_content
return [cast(AllMessageValues, first)] + list(messages[1:])
system_message: AllMessageValues = {"role": "system", "content": system_prompt}
return [system_message, *messages]
def convert_prefix_message_to_non_prefix_messages(
messages: List[AllMessageValues],
) -> List[AllMessageValues]:
+74 -37
View File
@@ -46,6 +46,7 @@ from litellm.types.llms.openai import (
ChatCompletionToolCallChunk,
ChatCompletionToolCallFunctionChunk,
ChatCompletionToolParam,
OpenAIChatCompletionFinishReason,
OpenAIMcpServerTool,
OpenAIWebSearchOptions,
)
@@ -54,10 +55,7 @@ from litellm.types.utils import (
CompletionTokensDetailsWrapper,
)
from litellm.types.utils import Message as LitellmMessage
from litellm.types.utils import (
PromptTokensDetailsWrapper,
ServerToolUse,
)
from litellm.types.utils import PromptTokensDetailsWrapper, ServerToolUse
from litellm.utils import (
ModelResponse,
Usage,
@@ -251,10 +249,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
# All numeric/string/array constraints not supported by Anthropic
unsupported_fields = {
"maxItems", "minItems", # array constraints
"minimum", "maximum", # numeric constraints
"exclusiveMinimum", "exclusiveMaximum", # numeric constraints
"minLength", "maxLength", # string constraints
"maxItems",
"minItems", # array constraints
"minimum",
"maximum", # numeric constraints
"exclusiveMinimum",
"exclusiveMaximum", # numeric constraints
"minLength",
"maxLength", # string constraints
}
# Build description additions from removed constraints
@@ -844,7 +846,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
@staticmethod
def map_openai_context_management_to_anthropic(
context_management: Union[List[Dict[str, Any]], Dict[str, Any]]
context_management: Union[List[Dict[str, Any]], Dict[str, Any]],
) -> Optional[Dict[str, Any]]:
"""
OpenAI format: [{"type": "compaction", "compact_threshold": 200000}]
@@ -876,19 +878,22 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
entry_type = entry.get("type")
if entry_type == "compaction":
anthropic_edit: Dict[str, Any] = {
"type": "compact_20260112"
}
anthropic_edit: Dict[str, Any] = {"type": "compact_20260112"}
compact_threshold = entry.get("compact_threshold")
# Rewrite to 'trigger' with correct nesting if threshold exists
if compact_threshold is not None and isinstance(compact_threshold, (int, float)):
if compact_threshold is not None and isinstance(
compact_threshold, (int, float)
):
anthropic_edit["trigger"] = {
"type": "input_tokens",
"value": int(compact_threshold)
"value": int(compact_threshold),
}
# Map any other keys by passthrough except handled ones
for k in entry:
if k not in {"type", "compact_threshold"}: # only passthrough other keys
if k not in {
"type",
"compact_threshold",
}: # only passthrough other keys
anthropic_edit[k] = entry[k]
anthropic_edits.append(anthropic_edit)
@@ -911,10 +916,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
for param, value in non_default_params.items():
if param == "max_tokens":
optional_params["max_tokens"] = value
if param == "max_completion_tokens":
optional_params["max_tokens"] = value
if param == "tools":
optional_params["max_tokens"] = (
value if isinstance(value, int) else max(1, int(round(value)))
)
elif param == "max_completion_tokens":
optional_params["max_tokens"] = (
value if isinstance(value, int) else max(1, int(round(value)))
)
elif param == "tools":
# check if optional params already has tools
anthropic_tools, mcp_servers = self._map_tools(value)
optional_params = self._add_tools_to_optional_params(
@@ -922,7 +931,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
)
if mcp_servers:
optional_params["mcp_servers"] = mcp_servers
if param == "tool_choice" or param == "parallel_tool_calls":
elif param == "tool_choice" or param == "parallel_tool_calls":
_tool_choice: Optional[AnthropicMessagesToolChoice] = (
self._map_tool_choice(
tool_choice=non_default_params.get("tool_choice"),
@@ -932,17 +941,19 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
if _tool_choice is not None:
optional_params["tool_choice"] = _tool_choice
if param == "stream" and value is True:
elif param == "stream" and value is True:
optional_params["stream"] = value
if param == "stop" and (isinstance(value, str) or isinstance(value, list)):
elif param == "stop" and (
isinstance(value, str) or isinstance(value, list)
):
_value = self._map_stop_sequences(value)
if _value is not None:
optional_params["stop_sequences"] = _value
if param == "temperature":
elif param == "temperature":
optional_params["temperature"] = value
if param == "top_p":
elif param == "top_p":
optional_params["top_p"] = value
if param == "response_format" and isinstance(value, dict):
elif param == "response_format" and isinstance(value, dict):
if any(
substring in model
for substring in {
@@ -982,14 +993,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
optional_params=optional_params, tools=[_tool]
)
optional_params["json_mode"] = True
if (
elif (
param == "user"
and value is not None
and isinstance(value, str)
and _valid_user_id(value) # anthropic fails on emails
):
optional_params["metadata"] = {"user_id": value}
if param == "thinking":
elif param == "thinking":
optional_params["thinking"] = value
elif param == "reasoning_effort" and isinstance(value, str):
optional_params["thinking"] = AnthropicConfig._map_reasoning_effort(
@@ -1007,9 +1018,13 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
elif param == "context_management":
# Supports both OpenAI list format and Anthropic dict format
if isinstance(value, (list, dict)):
anthropic_context_management = self.map_openai_context_management_to_anthropic(value)
anthropic_context_management = (
self.map_openai_context_management_to_anthropic(value)
)
if anthropic_context_management is not None:
optional_params["context_management"] = anthropic_context_management
optional_params["context_management"] = (
anthropic_context_management
)
elif param == "speed" and isinstance(value, str):
# Pass through Anthropic-specific speed parameter for fast mode
optional_params["speed"] = value
@@ -1071,7 +1086,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
if not system_message_block["content"]:
continue
# Skip system messages containing x-anthropic-billing-header metadata
if system_message_block["content"].startswith("x-anthropic-billing-header:"):
if system_message_block["content"].startswith(
"x-anthropic-billing-header:"
):
continue
anthropic_system_message_content = AnthropicSystemMessageContent(
type="text",
@@ -1091,7 +1108,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
if _content.get("type") == "text" and not text_value:
continue
# Skip system messages containing x-anthropic-billing-header metadata
if _content.get("type") == "text" and text_value and text_value.startswith("x-anthropic-billing-header:"):
if (
_content.get("type") == "text"
and text_value
and text_value.startswith("x-anthropic-billing-header:")
):
continue
anthropic_system_message_content = (
AnthropicSystemMessageContent(
@@ -1201,7 +1222,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
# Add context management header if any other edits/entries exist
if has_other:
self._ensure_beta_header(
headers, ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value
headers,
ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value,
)
def update_headers_with_optional_anthropic_beta(
@@ -1227,7 +1249,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
ANTHROPIC_HOSTED_TOOLS.MEMORY.value
):
self._ensure_beta_header(
headers, ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value
headers,
ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value,
)
if optional_params.get("context_management") is not None:
self._ensure_context_management_beta_header(
@@ -1491,7 +1514,16 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
if thinking_content is not None:
reasoning_content += thinking_content
return text_content, citations, thinking_blocks, reasoning_content, tool_calls, web_search_results, tool_results, compaction_blocks
return (
text_content,
citations,
thinking_blocks,
reasoning_content,
tool_calls,
web_search_results,
tool_results,
compaction_blocks,
)
def calculate_usage(
self,
@@ -1576,7 +1608,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
)
completion_token_details = CompletionTokensDetailsWrapper(
reasoning_tokens=reasoning_tokens if reasoning_tokens > 0 else 0,
text_tokens=completion_tokens - reasoning_tokens if reasoning_tokens > 0 else completion_tokens,
text_tokens=(
completion_tokens - reasoning_tokens
if reasoning_tokens > 0
else completion_tokens
),
)
total_tokens = prompt_tokens + completion_tokens
@@ -1696,8 +1732,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
"content"
] # allow user to access raw anthropic tool calling response
model_response.choices[0].finish_reason = map_finish_reason(
completion_response["stop_reason"]
model_response.choices[0].finish_reason = cast(
OpenAIChatCompletionFinishReason,
map_finish_reason(completion_response["stop_reason"]),
)
## CALCULATING USAGE
+42 -18
View File
@@ -159,6 +159,7 @@ from .litellm_core_utils.fallback_utils import (
completion_with_fallbacks,
)
from .litellm_core_utils.prompt_templates.common_utils import (
add_system_prompt_to_messages,
get_completion_messages,
update_messages_with_model_file_ids,
)
@@ -599,7 +600,7 @@ async def acompletion( # noqa: PLR0915
# Add the context to the function
ctx = contextvars.copy_context()
func_with_context = partial(ctx.run, func)
init_response = await loop.run_in_executor(None, func_with_context)
if isinstance(init_response, dict) or isinstance(
init_response, ModelResponse
@@ -939,7 +940,7 @@ def responses_api_bridge_check(
model = model.replace("responses/", "")
mode = "responses"
model_info["mode"] = mode
if web_search_options is not None and custom_llm_provider == "xai":
model_info["mode"] = "responses"
model = model.replace("responses/", "")
@@ -1108,9 +1109,7 @@ def completion( # type: ignore # noqa: PLR0915
skip_mcp_handler = kwargs.pop("_skip_mcp_handler", False)
if not skip_mcp_handler and tools:
from litellm.responses.mcp.chat_completions_handler import (
acompletion_with_mcp,
)
from litellm.responses.mcp.chat_completions_handler import acompletion_with_mcp
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
LiteLLM_Proxy_MCP_Handler,
)
@@ -1245,6 +1244,7 @@ def completion( # type: ignore # noqa: PLR0915
### PROMPT MANAGEMENT ###
prompt_id = cast(Optional[str], kwargs.get("prompt_id", None))
prompt_variables = cast(Optional[dict], kwargs.get("prompt_variables", None))
litellm_system_prompt = kwargs.get("litellm_system_prompt", None)
### COPY MESSAGES ### - related issue https://github.com/BerriAI/litellm/discussions/4489
messages = get_completion_messages(
messages=messages,
@@ -1276,6 +1276,14 @@ def completion( # type: ignore # noqa: PLR0915
prompt_version=kwargs.get("prompt_version", None),
)
### LITELLM SYSTEM PROMPT ###
if litellm_system_prompt:
messages = add_system_prompt_to_messages(
messages=messages,
system_prompt=litellm_system_prompt,
merge_with_first_system=True,
)
try:
if base_url is not None:
api_base = base_url
@@ -1558,7 +1566,9 @@ def completion( # type: ignore # noqa: PLR0915
## RESPONSES API BRIDGE LOGIC ## - check if model has 'mode: responses' in litellm.model_cost map
model_info, model = responses_api_bridge_check(
model=model, custom_llm_provider=custom_llm_provider, web_search_options=web_search_options
model=model,
custom_llm_provider=custom_llm_provider,
web_search_options=web_search_options,
)
if model_info.get("mode") == "responses":
@@ -2209,17 +2219,19 @@ def completion( # type: ignore # noqa: PLR0915
elif custom_llm_provider == "a2a":
# A2A (Agent-to-Agent) Protocol
# Resolve agent configuration from registry if model format is "a2a/<agent-name>"
api_base, api_key, headers = litellm.A2AConfig.resolve_agent_config_from_registry(
model=model,
api_base=api_base,
api_key=api_key,
headers=headers,
optional_params=optional_params,
api_base, api_key, headers = (
litellm.A2AConfig.resolve_agent_config_from_registry(
model=model,
api_base=api_base,
api_key=api_key,
headers=headers,
optional_params=optional_params,
)
)
# Fall back to environment variables and defaults
api_base = api_base or litellm.api_base or get_secret_str("A2A_API_BASE")
if api_base is None:
raise Exception(
"api_base is required for A2A provider. "
@@ -4783,7 +4795,10 @@ def embedding( # noqa: PLR0915
or custom_llm_provider == "together_ai"
or custom_llm_provider == "nvidia_nim"
or custom_llm_provider == "litellm_proxy"
or (model in litellm.open_ai_embedding_models and custom_llm_provider is None)
or (
model in litellm.open_ai_embedding_models
and custom_llm_provider is None
)
):
api_base = (
api_base
@@ -7239,7 +7254,11 @@ def stream_chunk_builder( # noqa: PLR0915
continue
choice = chunk["choices"][0]
delta_obj = choice.get("delta", {}) if isinstance(choice, dict) else getattr(choice, "delta", {})
delta_obj = (
choice.get("delta", {})
if isinstance(choice, dict)
else getattr(choice, "delta", {})
)
if isinstance(delta_obj, dict):
delta = delta_obj
elif hasattr(delta_obj, "model_dump"):
@@ -7266,7 +7285,9 @@ def stream_chunk_builder( # noqa: PLR0915
if is_simple_text_stream:
if simple_content_parts:
response["choices"][0]["message"]["content"] = "".join(simple_content_parts)
response["choices"][0]["message"]["content"] = "".join(
simple_content_parts
)
completion_output = get_content_from_model_response(response)
usage = processor.calculate_usage(
chunks=chunks,
@@ -7291,7 +7312,9 @@ def stream_chunk_builder( # noqa: PLR0915
if litellm.include_cost_in_streaming_usage and logging_obj is not None:
setattr(
usage, "cost", logging_obj._response_cost_calculator(result=response)
usage,
"cost",
logging_obj._response_cost_calculator(result=response),
)
return response
@@ -7504,6 +7527,7 @@ def __getattr__(name: str) -> Any:
# before loading tiktoken, ensuring the local cache is used
# instead of downloading from the internet
from litellm._lazy_imports import _get_default_encoding
_encoding = _get_default_encoding()
# Cache it in the module's __dict__ for subsequent accesses
import sys
+4
View File
@@ -16,6 +16,10 @@ model_list:
- model_name: gpt-5-mini
litellm_params:
model: openai/gpt-5-mini
- model_name: custom_litellm_model
litellm_params:
model: litellm_agent/claude-sonnet-4-5-20250929
litellm_system_prompt: "Be a helpful assistant."
guardrails:
@@ -2,6 +2,9 @@
This module allows users to write custom guardrail logic using Python-like code
that runs in a sandboxed environment with access to LiteLLM-provided primitives.
Pre-built custom code for common guardrails (e.g. response rejection detection)
is available in response_rejection_code.py.
"""
from typing import TYPE_CHECKING
@@ -9,6 +12,8 @@ from typing import TYPE_CHECKING
from litellm.types.guardrails import SupportedGuardrailIntegrations
from .custom_code_guardrail import CustomCodeGuardrail
from .response_rejection_code import (DEFAULT_REJECTION_PHRASES,
RESPONSE_REJECTION_GUARDRAIL_CODE)
if TYPE_CHECKING:
from litellm.types.guardrails import Guardrail, LitellmParams
@@ -61,5 +66,7 @@ guardrail_class_registry = {
__all__ = [
"CustomCodeGuardrail",
"DEFAULT_REJECTION_PHRASES",
"RESPONSE_REJECTION_GUARDRAIL_CODE",
"initialize_guardrail",
]
@@ -26,6 +26,12 @@ Example custom code (async with HTTP):
if response["success"] and response["body"].get("flagged"):
return block("Content flagged by moderation API")
return allow()
Example: block when response rejects the user (input_type response only):
Use RESPONSE_REJECTION_GUARDRAIL_CODE from .response_rejection_code it
checks response texts for phrases like "That's not something I can help with"
and returns block() so the guardrail raises a block error.
"""
import asyncio
@@ -35,18 +41,18 @@ from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Type, cast
from fastapi import HTTPException
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.integrations.custom_guardrail import (CustomGuardrail,
log_guardrail_information)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
from litellm.types.proxy.guardrails.guardrail_hooks.base import \
GuardrailConfigModel
from litellm.types.utils import GenericGuardrailAPIInputs
from .primitives import get_custom_code_primitives
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import \
Logging as LiteLLMLoggingObj
class CustomCodeGuardrailError(Exception):
@@ -0,0 +1,76 @@
"""
Custom code for a response guardrail that blocks when the model response
indicates it is rejecting the user request (e.g. "That's not something I can help with").
Use this with the Custom Code Guardrail (custom_code) by setting litellm_params.custom_code
to RESPONSE_REJECTION_GUARDRAIL_CODE. The guardrail runs only on input_type "response"
and raises a block error if any response text matches known rejection phrases.
"""
# Default phrases that indicate the model is refusing the user request (lowercase for case-insensitive match).
# Custom code guardrails can override by defining rejection_phrases in the code.
DEFAULT_REJECTION_PHRASES = [
"that's not something i can help with",
"that is not something i can help with",
"i can't help with that",
"i cannot help with that",
"i'm not able to help",
"i am not able to help",
"i'm unable to help",
"i cannot assist",
"i can't assist",
"i'm not allowed to",
"i'm not permitted to",
"i won't be able to help",
"i'm sorry, i can't",
"i'm sorry, i cannot",
"as an ai, i can't",
"as an ai, i cannot",
]
# Custom code string for the Custom Code Guardrail. Only runs on input_type "response".
# Uses primitives: allow(), block(), lower(), contains()
RESPONSE_REJECTION_GUARDRAIL_CODE = '''
def apply_guardrail(inputs, request_data, input_type):
"""Block responses that indicate the model rejected the user request."""
if input_type != "response":
return allow()
texts = inputs.get("texts") or []
# All lowercase for case-insensitive matching (text is lowercased before check)
rejection_phrases = [
"that's not something i can help with",
"that is not something i can help with",
"i can't help with that",
"i cannot help with that",
"i'm not able to help",
"i am not able to help",
"i'm unable to help",
"i cannot assist",
"i can't assist",
"i'm not allowed to",
"i'm not permitted to",
"i won't be able to help",
"i'm sorry, i can't",
"i'm sorry, i cannot",
"as an ai, i can't",
"as an ai, i cannot",
]
for text in texts:
if not text:
continue
text_lower = lower(text)
for phrase in rejection_phrases:
if contains(text_lower, phrase):
return block(
"Response indicates the model rejected the user request.",
detection_info={"matched_phrase": phrase, "input_type": "response"},
)
return allow()
'''
__all__ = [
"DEFAULT_REJECTION_PHRASES",
"RESPONSE_REJECTION_GUARDRAIL_CODE",
]
@@ -20,7 +20,6 @@ from datetime import datetime, timedelta, timezone
from typing import Any, Dict, List, Literal, Optional, Tuple, cast
import fastapi
import prisma
import yaml
from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, status
@@ -1072,6 +1071,7 @@ async def generate_key_fn(
- team_id: Optional[str] - The team id of the key
- user_id: Optional[str] - The user id of the key
- organization_id: Optional[str] - The organization id of the key. If not set, and team_id is set, the organization id will be the same as the team id. If conflict, an error will be raised.
- project_id: Optional[str] - The project id of the key. When set, models and max_budget are validated against the project's limits.
- budget_id: Optional[str] - The budget id associated with the key. Created by calling `/budget/new`.
- models: Optional[list] - Model_name's a user is allowed to call. (if empty, key is allowed to call all models)
- aliases: Optional[dict] - Any alias mappings, on top of anything in the config.yaml model list. - https://docs.litellm.ai/docs/proxy/virtual_keys#managing-auth---upgradedowngrade-models
@@ -3153,6 +3153,8 @@ async def _rotate_master_key( # noqa: PLR0915
3. Encrypt the values with the new master key
4. Update the values in the DB
"""
import prisma
from litellm.proxy.proxy_server import proxy_config
try:
@@ -12,10 +12,11 @@ All /policy management endpoints
import copy
import json
import os
from typing import TYPE_CHECKING, AsyncIterator, List, Literal, Optional, cast
from typing import (TYPE_CHECKING, Any, AsyncIterator, List, Literal, Optional,
cast)
from fastapi import APIRouter, Depends, HTTPException, Request
from fastapi.responses import StreamingResponse
from fastapi.responses import Response, StreamingResponse
from pydantic import BaseModel, Field
from typing_extensions import TypedDict
@@ -24,8 +25,12 @@ from litellm.constants import (COMPETITOR_LLM_TEMPERATURE,
DEFAULT_COMPETITOR_DISCOVERY_MODEL,
MAX_COMPETITOR_NAMES)
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.llms.openai.chat.guardrail_translation.handler import \
OpenAIChatCompletionsHandler
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.guardrails.guardrail_hooks.custom_code import (
RESPONSE_REJECTION_GUARDRAIL_CODE, CustomCodeGuardrail)
from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry
from litellm.proxy.management_helpers.utils import management_endpoint_wrapper
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
@@ -39,7 +44,7 @@ from litellm.types.proxy.policy_engine import (PolicyGuardrailsResponse,
PolicyTestResponse,
PolicyValidateRequest,
PolicyValidationResponse)
from litellm.types.utils import GenericGuardrailAPIInputs
from litellm.types.utils import GenericGuardrailAPIInputs, ModelResponse
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import \
@@ -69,20 +74,32 @@ class GuardrailErrorEntry(TypedDict):
message: str
class ApplyPoliciesResult(TypedDict):
"""Result of apply_policies: inputs plus any guardrail failures."""
class _ApplyPoliciesResultBase(TypedDict):
"""Base result of apply_policies: inputs plus any guardrail failures."""
inputs: GenericGuardrailAPIInputs
guardrail_errors: List[GuardrailErrorEntry]
class ApplyPoliciesPerItemResult(TypedDict):
"""Result for one input when using inputs_list."""
class ApplyPoliciesResult(_ApplyPoliciesResultBase, total=False):
"""Result of apply_policies. agent_response set when agent_id provided."""
agent_response: Any
class _ApplyPoliciesPerItemResultBase(TypedDict):
"""Base result for one input when using inputs_list."""
inputs: GenericGuardrailAPIInputs
guardrail_errors: List[GuardrailErrorEntry]
class ApplyPoliciesPerItemResult(_ApplyPoliciesPerItemResultBase, total=False):
"""Result for one input when using inputs_list. agent_response set when agent_id provided."""
agent_response: Any
class ApplyPoliciesListResult(TypedDict):
"""Result when using inputs_list: one result per input."""
@@ -185,21 +202,78 @@ async def apply_policies(
return {"inputs": current_inputs, "guardrail_errors": guardrail_errors}
def _chat_body_from_inputs(
inputs: GenericGuardrailAPIInputs, agent_id: str, request_data: dict
) -> dict:
"""Build a chat completion request body from guardrail inputs and agent_id."""
messages: List[dict]
structured = inputs.get("structured_messages")
texts = inputs.get("texts")
if structured:
messages = list(structured) # type: ignore[arg-type]
elif texts:
if len(texts) == 1:
messages = [{"role": "user", "content": texts[0]}]
else:
messages = [{"role": "user", "content": "\n".join(texts)}]
else:
messages = [{"role": "user", "content": "Hello"}]
body: dict = {"model": agent_id, "messages": messages, "stream": False}
if request_data:
body.setdefault("metadata", {}).update(request_data)
return body
def _request_with_json_body(body: dict) -> Request:
"""Create a Starlette Request that will return the given dict as parsed JSON body."""
body_bytes = json.dumps(body).encode()
received: List[bool] = [False]
async def receive() -> dict:
if received[0]:
return {"type": "http.disconnect"}
received[0] = True
return {"type": "http.request", "body": body_bytes, "more_body": False}
scope: dict = {
"type": "http",
"method": "POST",
"path": "/v1/chat/completions",
"query_string": b"",
"headers": [(b"content-type", b"application/json")],
"scheme": "http",
"server": ("localhost", 8000),
"client": ("127.0.0.1", 0),
"root_path": "",
"app": None,
"asgi": {"version": "3.0", "spec_version": "2.0"},
}
return Request(scope, receive=receive)
class TestPoliciesAndGuardrailsRequest(BaseModel):
"""Request body for POST /utils/test_policies_and_guardrails."""
policy_names: Optional[List[str]] = Field(default=None, description="Policy names to resolve guardrails from")
guardrail_names: Optional[List[str]] = Field(default=None, description="Guardrail names to apply directly")
inputs: Optional[dict] = Field(
default=None,
description="GenericGuardrailAPIInputs, e.g. { \"texts\": [\"...\"] }. Use inputs_list for per-input processing.",
policy_names: Optional[List[str]] = Field(
default=None, description="Policy names to resolve guardrails from"
)
inputs_list: Optional[List[dict]] = Field(
default=None,
guardrail_names: Optional[List[str]] = Field(
default=None, description="Guardrail names to apply directly"
)
inputs_list: List[GenericGuardrailAPIInputs] = Field(
default=[],
description="List of GenericGuardrailAPIInputs; each item processed separately (for batch compliance testing).",
)
request_data: dict = Field(default_factory=dict, description="Request context (model, user_id, etc.)")
input_type: Literal["request", "response"] = Field(default="request", description="Whether inputs are request or response")
request_data: dict = Field(
default_factory=dict, description="Request context (model, user_id, etc.)"
)
input_type: Literal["request", "response"] = Field(
default="request", description="Whether inputs are request or response"
)
agent_id: Optional[str] = Field(
default=None,
description="When set, call chat completion with this model/agent for each input and include the response in the result.",
)
@router.post(
@@ -223,40 +297,86 @@ async def test_policies_and_guardrails(
"""
from litellm.litellm_core_utils.litellm_logging import \
Logging as LiteLLMLoggingObj
from litellm.proxy.proxy_server import proxy_logging_obj
from litellm.proxy.proxy_server import chat_completion, proxy_logging_obj
from litellm.proxy.utils import handle_exception_on_proxy
def _serialize_chat_response(response: Any) -> Any:
if hasattr(response, "model_dump"):
return response.model_dump(exclude_unset=True)
if isinstance(response, dict):
return response
return response
async def _get_agent_response(
inputs: GenericGuardrailAPIInputs,
agent_id: str,
user_api_key_dict: UserAPIKeyAuth,
) -> Any:
body = _chat_body_from_inputs(inputs, agent_id, data.request_data)
req = _request_with_json_body(body)
resp = Response()
result = await chat_completion(
request=req,
fastapi_response=resp,
model=agent_id,
user_api_key_dict=user_api_key_dict,
)
return _serialize_chat_response(result)
try:
logging_obj = cast(LiteLLMLoggingObj, proxy_logging_obj)
if data.inputs_list is not None:
results: List[ApplyPoliciesPerItemResult] = []
for inp in data.inputs_list:
inputs_typed = cast(GenericGuardrailAPIInputs, inp)
item_result = await apply_policies(
policy_names=data.policy_names,
inputs=inputs_typed,
request_data=data.request_data,
input_type=data.input_type,
proxy_logging_obj=logging_obj,
guardrail_names=data.guardrail_names,
)
results.append(
ApplyPoliciesPerItemResult(
inputs=item_result["inputs"],
guardrail_errors=item_result["guardrail_errors"],
)
)
return ApplyPoliciesListResult(results=results)
if data.inputs is not None:
inputs_typed = cast(GenericGuardrailAPIInputs, data.inputs)
return await apply_policies(
results: List[ApplyPoliciesPerItemResult] = []
for inp in data.inputs_list:
item_result = await apply_policies(
policy_names=data.policy_names,
inputs=inputs_typed,
inputs=inp,
request_data=data.request_data,
input_type=data.input_type,
proxy_logging_obj=logging_obj,
guardrail_names=data.guardrail_names,
)
item: ApplyPoliciesPerItemResult = {
"inputs": item_result["inputs"],
"guardrail_errors": item_result["guardrail_errors"],
}
if data.agent_id is not None:
item["agent_response"] = await _get_agent_response(
item_result["inputs"],
data.agent_id,
user_api_key_dict,
)
# run response through response_rejection_guardrail (reuses handler extraction + apply)
response_rejection_guardrail = CustomCodeGuardrail(
custom_code=RESPONSE_REJECTION_GUARDRAIL_CODE,
guardrail_name="response_rejection",
)
try:
model_response = ModelResponse.model_validate(
item["agent_response"]
)
handler = OpenAIChatCompletionsHandler()
await handler.process_output_response(
response=model_response,
guardrail_to_apply=response_rejection_guardrail,
litellm_logging_obj=logging_obj,
user_api_key_dict=user_api_key_dict,
)
except Exception as guardrail_err:
item["guardrail_errors"] = list(item["guardrail_errors"])
detail = getattr(guardrail_err, "detail", None)
if isinstance(detail, dict) and "error" in detail:
err_msg = detail["error"]
else:
err_msg = str(detail if detail is not None else guardrail_err)
item["guardrail_errors"].append(
GuardrailErrorEntry(
guardrail_name="response_rejection",
message=err_msg,
)
)
results.append(item)
return ApplyPoliciesListResult(results=results)
raise ValueError("Either inputs or inputs_list must be provided")
except Exception as e:
raise handle_exception_on_proxy(e)
@@ -581,11 +701,15 @@ def _validate_enrichment_request(data: EnrichTemplateRequest) -> tuple[dict, dic
templates = _load_policy_templates_from_local_backup()
template = next((t for t in templates if t.get("id") == data.template_id), None)
if template is None:
raise HTTPException(status_code=404, detail=f"Template '{data.template_id}' not found")
raise HTTPException(
status_code=404, detail=f"Template '{data.template_id}' not found"
)
llm_enrichment = template.get("llm_enrichment")
if llm_enrichment is None:
raise HTTPException(status_code=400, detail="Template does not support LLM enrichment")
raise HTTPException(
status_code=400, detail="Template does not support LLM enrichment"
)
# Validate competitors list size if provided
if data.competitors and len(data.competitors) > MAX_COMPETITOR_NAMES:
@@ -695,7 +819,11 @@ async def _stream_llm_competitor_names(
while "\n" in buffer:
line, buffer = buffer.split("\n", 1)
name = _clean_competitor_line(line)
if name and name.lower() not in existing_lower and count < MAX_COMPETITOR_NAMES:
if (
name
and name.lower() not in existing_lower
and count < MAX_COMPETITOR_NAMES
):
existing_lower.add(name.lower())
count += 1
yield name, False
@@ -744,9 +872,7 @@ async def _stream_competitor_events(
"{{" + llm_enrichment["parameter"] + "}}", brand_name
)
try:
async for name, _ in _stream_llm_competitor_names(
prompt, model, []
):
async for name, _ in _stream_llm_competitor_names(prompt, model, []):
if name:
competitors.append(name)
yield f"data: {json.dumps({'type': 'competitor', 'name': name})}\n\n"
@@ -893,10 +1019,7 @@ def _build_all_names_per_competitor(
competitors: list[str], variations_map: dict[str, list[str]]
) -> dict[str, list[str]]:
"""Build canonical + variation name lists for each competitor."""
return {
comp: [comp] + variations_map.get(comp, [])
for comp in competitors
}
return {comp: [comp] + variations_map.get(comp, []) for comp in competitors}
def _build_competitor_guardrail_definitions(
@@ -912,7 +1035,9 @@ def _build_competitor_guardrail_definitions(
output_blocked = _build_name_blocked_words(competitors, all_names)
recommendation_blocked = _build_recommendation_blocked_words(competitors, all_names)
comparison_blocked = _build_comparison_blocked_words(competitors, all_names, brand_name)
comparison_blocked = _build_comparison_blocked_words(
competitors, all_names, brand_name
)
blocked_words_map = {
"competitor-output-blocker": output_blocked,
@@ -943,7 +1068,11 @@ def _build_name_blocked_words(
result = []
for comp in competitors:
for name in all_names[comp]:
desc = f"Competitor: {comp}" if name == comp else f"Competitor variation ({comp}): {name}"
desc = (
f"Competitor: {comp}"
if name == comp
else f"Competitor variation ({comp}): {name}"
)
result.append({"keyword": name, "action": "BLOCK", "description": desc})
return result
@@ -956,11 +1085,13 @@ def _build_recommendation_blocked_words(
for comp in competitors:
for name in all_names[comp]:
for prefix in ["try", "use", "switch to", "consider"]:
result.append({
"keyword": f"{prefix} {name}",
"action": "BLOCK",
"description": f"Recommendation to competitor ({comp})",
})
result.append(
{
"keyword": f"{prefix} {name}",
"action": "BLOCK",
"description": f"Recommendation to competitor ({comp})",
}
)
return result
@@ -971,23 +1102,29 @@ def _build_comparison_blocked_words(
result = []
for comp in competitors:
for name in all_names[comp]:
result.append({
"keyword": f"{name} is better",
"action": "BLOCK",
"description": f"Unfavorable comparison ({comp})",
})
result.append(
{
"keyword": f"{name} is better",
"action": "BLOCK",
"description": f"Unfavorable comparison ({comp})",
}
)
# Brand-level comparisons (only need one entry each, not per-competitor)
result.append({
"keyword": f"better than {brand_name}",
"action": "BLOCK",
"description": "Unfavorable comparison",
})
result.append({
"keyword": f"{brand_name} is worse",
"action": "BLOCK",
"description": "Unfavorable comparison",
})
result.append(
{
"keyword": f"better than {brand_name}",
"action": "BLOCK",
"description": "Unfavorable comparison",
}
)
result.append(
{
"keyword": f"{brand_name} is worse",
"action": "BLOCK",
"description": "Unfavorable comparison",
}
)
return result
@@ -1121,7 +1258,9 @@ async def _test_guardrail_definitions(
request_data={},
input_type="request",
)
output_text = output.get("texts", [text])[0] if output.get("texts") else text
output_text = (
output.get("texts", [text])[0] if output.get("texts") else text
)
if output_text != text:
action = "masked"
+15 -12
View File
@@ -163,11 +163,7 @@ from litellm.types.utils import (
)
from litellm.types.utils import ModelInfo
from litellm.types.utils import ModelInfo as ModelMapInfo
from litellm.types.utils import (
ModelResponseStream,
StandardLoggingPayload,
Usage,
)
from litellm.types.utils import ModelResponseStream, StandardLoggingPayload, Usage
from litellm.utils import (
CustomStreamWrapper,
EmbeddingResponse,
@@ -1996,12 +1992,11 @@ class Router:
When both have tools, concatenate them (deployment tools first, then request tools).
tool_choice: use request value if provided, else deployment's.
"""
dep_params = deployment.get("litellm_params", {}) or {}
dep_params = (
dep_params.model_dump(exclude_none=True)
if hasattr(dep_params, "model_dump")
else dep_params
)
dep_params_raw = deployment.get("litellm_params", {}) or {}
if isinstance(dep_params_raw, dict):
dep_params = dep_params_raw
else:
dep_params = dep_params_raw.model_dump(exclude_none=True)
dep_tools = dep_params.get("tools") or []
req_tools = kwargs.get("tools") or []
if dep_tools or req_tools:
@@ -2573,6 +2568,12 @@ class Router:
litellm_model = data.get("model", None)
# litellm_agent/ prefix only strips the model name, no prompt_id needed
is_litellm_agent_model = (
isinstance(litellm_model, str)
and litellm_model.startswith("litellm_agent/")
)
prompt_id = kwargs.get("prompt_id") or prompt_management_deployment[
"litellm_params"
].get("prompt_id", None)
@@ -2585,7 +2586,9 @@ class Router:
"litellm_params"
].get("prompt_label", None)
if prompt_id is None or not isinstance(prompt_id, str):
if not is_litellm_agent_model and (
prompt_id is None or not isinstance(prompt_id, str)
):
raise ValueError(
f"Prompt ID is not set or not a string. Got={prompt_id}, type={type(prompt_id)}"
)
+29 -40
View File
@@ -1,61 +1,47 @@
import json
import time
from enum import Enum
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Mapping, Optional, Union
from typing import (TYPE_CHECKING, Any, Dict, List, Literal, Mapping, Optional,
Union)
from openai._models import BaseModel as OpenAIObject
from openai.types.audio.transcription_create_params import (
FileTypes as FileTypes, # type: ignore
)
from openai.types.audio.transcription_create_params import \
FileTypes as FileTypes # type: ignore
from openai.types.chat.chat_completion import ChatCompletion as ChatCompletion
from openai.types.completion_usage import (
CompletionTokensDetails,
CompletionUsage,
PromptTokensDetails,
)
from openai.types.completion_usage import (CompletionTokensDetails,
CompletionUsage,
PromptTokensDetails)
from openai.types.moderation import Categories as Categories
from openai.types.moderation import (
CategoryAppliedInputTypes as CategoryAppliedInputTypes,
)
from openai.types.moderation import \
CategoryAppliedInputTypes as CategoryAppliedInputTypes
from openai.types.moderation import CategoryScores as CategoryScores
from openai.types.moderation_create_response import Moderation as Moderation
from openai.types.moderation_create_response import (
ModerationCreateResponse as ModerationCreateResponse,
)
from openai.types.moderation_create_response import \
ModerationCreateResponse as ModerationCreateResponse
from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, model_validator
from typing_extensions import Required, TypedDict
from litellm._uuid import uuid
from litellm.types.llms.base import (
BaseLiteLLMOpenAIResponseObject,
LiteLLMPydanticObjectBase,
)
from litellm.types.llms.base import (BaseLiteLLMOpenAIResponseObject,
LiteLLMPydanticObjectBase)
from litellm.types.mcp import MCPServerCostInfo
from ..litellm_core_utils.core_helpers import map_finish_reason
from .agents import LiteLLMSendMessageResponse
from .guardrails import GuardrailEventHooks
from .llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse
from .llms.anthropic_messages.anthropic_response import \
AnthropicMessagesResponse
from .llms.base import HiddenParams
from .llms.openai import (
AllMessageValues,
Batch,
ChatCompletionAnnotation,
ChatCompletionRedactedThinkingBlock,
ChatCompletionThinkingBlock,
ChatCompletionToolCallChunk,
ChatCompletionToolParam,
ChatCompletionUsageBlock,
FileSearchTool,
FineTuningJob,
ImageURLListItem,
OpenAIChatCompletionChunk,
OpenAIChatCompletionFinishReason,
OpenAIFileObject,
OpenAIRealtimeStreamList,
ResponsesAPIResponse,
WebSearchOptions,
)
from .llms.openai import (AllMessageValues, Batch, ChatCompletionAnnotation,
ChatCompletionRedactedThinkingBlock,
ChatCompletionThinkingBlock,
ChatCompletionToolCallChunk, ChatCompletionToolParam,
ChatCompletionUsageBlock, FileSearchTool,
FineTuningJob, ImageURLListItem,
OpenAIChatCompletionChunk,
OpenAIChatCompletionFinishReason, OpenAIFileObject,
OpenAIRealtimeStreamList, ResponsesAPIResponse,
WebSearchOptions)
from .rerank import RerankResponse as RerankResponse
if TYPE_CHECKING:
@@ -2917,8 +2903,9 @@ all_litellm_params = (
"api_key",
"api_version",
"prompt_id",
"provider_specific_header",
"prompt_variables",
"litellm_system_prompt",
"provider_specific_header",
"prompt_version",
"api_base",
"force_timeout",
@@ -3172,6 +3159,7 @@ class LlmProviders(str, Enum):
POE = "poe"
CHUTES = "chutes"
XIAOMI_MIMO = "xiaomi_mimo"
LITELLM_AGENT = "litellm_agent"
# Create a set of all provider values for quick lookup
@@ -3202,6 +3190,7 @@ class SearchProviders(str, Enum):
LINKUP = "linkup"
DUCKDUCKGO = "duckduckgo"
# Create a set of all search provider values for quick lookup
SearchProvidersSet = {provider.value for provider in SearchProviders}
Generated
+56 -44
View File
@@ -1,4 +1,4 @@
# This file is automatically @generated by Poetry 2.3.2 and should not be changed by hand.
# This file is automatically @generated by Poetry 2.2.0 and should not be changed by hand.
[[package]]
name = "a2a-sdk"
@@ -7,11 +7,11 @@ description = "A2A Python SDK"
optional = false
python-versions = ">=3.10"
groups = ["main", "proxy-dev"]
markers = "python_version >= \"3.10\""
files = [
{file = "a2a_sdk-0.3.22-py3-none-any.whl", hash = "sha256:b98701135bb90b0ff85d35f31533b6b7a299bf810658c1c65f3814a6c15ea385"},
{file = "a2a_sdk-0.3.22.tar.gz", hash = "sha256:77a5694bfc4f26679c11b70c7f1062522206d430b34bc1215cfbb1eba67b7e7d"},
]
markers = {main = "python_version >= \"3.10\" and extra == \"extra-proxy\"", proxy-dev = "python_version >= \"3.10\""}
[package.dependencies]
google-api-core = ">=1.26.0"
@@ -385,7 +385,6 @@ files = [
{file = "azure_core-1.36.0-py3-none-any.whl", hash = "sha256:fee9923a3a753e94a259563429f3644aaf05c486d45b1215d098115102d91d3b"},
{file = "azure_core-1.36.0.tar.gz", hash = "sha256:22e5605e6d0bf1d229726af56d9e92bc37b6e726b141a18be0b4d424131741b7"},
]
markers = {main = "extra == \"proxy\" or extra == \"extra-proxy\""}
[package.dependencies]
requests = ">=2.21.0"
@@ -406,7 +405,6 @@ files = [
{file = "azure_identity-1.25.1-py3-none-any.whl", hash = "sha256:e9edd720af03dff020223cd269fa3a61e8f345ea75443858273bcb44844ab651"},
{file = "azure_identity-1.25.1.tar.gz", hash = "sha256:87ca8328883de6036443e1c37b40e8dc8fb74898240f61071e09d2e369361456"},
]
markers = {main = "extra == \"proxy\" or extra == \"extra-proxy\""}
[package.dependencies]
azure-core = ">=1.31.0"
@@ -600,7 +598,7 @@ files = [
{file = "cachetools-6.2.2-py3-none-any.whl", hash = "sha256:6c09c98183bf58560c97b2abfcedcbaf6a896a490f534b031b661d3723b45ace"},
{file = "cachetools-6.2.2.tar.gz", hash = "sha256:8e6d266b25e539df852251cfd6f990b4bc3a141db73b939058d809ebd2590fc6"},
]
markers = {main = "python_version >= \"3.10\" and (extra == \"extra-proxy\" or extra == \"google\" or extra == \"mlflow\") or extra == \"google\" or extra == \"extra-proxy\"", proxy-dev = "python_version >= \"3.10\""}
markers = {main = "extra == \"google\" or extra == \"extra-proxy\" or python_version >= \"3.10\"", proxy-dev = "python_version >= \"3.10\""}
[[package]]
name = "certifi"
@@ -707,7 +705,7 @@ files = [
{file = "cffi-2.0.0-cp39-cp39-win_amd64.whl", hash = "sha256:b882b3df248017dba09d6b16defe9b5c407fe32fc7c65a9c69798e6175601be9"},
{file = "cffi-2.0.0.tar.gz", hash = "sha256:44d1b5909021139fe36001ae048dbdde8214afa20200eda0f64c068cac5d5529"},
]
markers = {main = "(platform_python_implementation != \"PyPy\" or extra == \"proxy\") and (python_version >= \"3.10\" or extra == \"proxy\" or extra == \"extra-proxy\") and (extra == \"proxy\" or extra == \"extra-proxy\" or extra == \"mlflow\")", dev = "platform_python_implementation != \"PyPy\"", proxy-dev = "platform_python_implementation != \"PyPy\""}
markers = {main = "platform_python_implementation != \"PyPy\" or extra == \"proxy\"", dev = "platform_python_implementation != \"PyPy\"", proxy-dev = "platform_python_implementation != \"PyPy\""}
[package.dependencies]
pycparser = {version = "*", markers = "implementation_name != \"PyPy\""}
@@ -1057,7 +1055,6 @@ files = [
{file = "cryptography-43.0.3-pp39-pypy39_pp73-win_amd64.whl", hash = "sha256:2ce6fae5bdad59577b44e4dfed356944fbf1d925269114c28be377692643b4ff"},
{file = "cryptography-43.0.3.tar.gz", hash = "sha256:315b9001266a492a6ff443b61238f956b214dbec9910a081ba5b6646a055a805"},
]
markers = {main = "python_version >= \"3.10\" and (extra == \"proxy\" or extra == \"extra-proxy\" or extra == \"mlflow\") or extra == \"proxy\" or extra == \"extra-proxy\""}
[package.dependencies]
cffi = {version = ">=1.12", markers = "platform_python_implementation != \"PyPy\""}
@@ -1840,11 +1837,11 @@ description = "Google API client core library"
optional = false
python-versions = ">=3.7"
groups = ["main", "proxy-dev"]
markers = "python_version >= \"3.14\""
files = [
{file = "google_api_core-2.25.2-py3-none-any.whl", hash = "sha256:e9a8f62d363dc8424a8497f4c2a47d6bcda6c16514c935629c257ab5d10210e7"},
{file = "google_api_core-2.25.2.tar.gz", hash = "sha256:1c63aa6af0d0d5e37966f157a77f9396d820fba59f9e43e9415bc3dc5baff300"},
]
markers = {main = "python_version >= \"3.14\" and (extra == \"extra-proxy\" or extra == \"google\")", proxy-dev = "python_version >= \"3.14\""}
[package.dependencies]
google-auth = ">=2.14.1,<3.0.0"
@@ -1872,7 +1869,7 @@ files = [
{file = "google_api_core-2.28.1-py3-none-any.whl", hash = "sha256:4021b0f8ceb77a6fb4de6fde4502cecab45062e66ff4f2895169e0b35bc9466c"},
{file = "google_api_core-2.28.1.tar.gz", hash = "sha256:2b405df02d68e68ce0fbc138559e6036559e685159d148ae5861013dc201baf8"},
]
markers = {main = "python_version < \"3.14\" and (extra == \"extra-proxy\" or extra == \"google\")", proxy-dev = "python_version >= \"3.10\" and python_version < \"3.14\""}
markers = {main = "(python_version >= \"3.10\" or extra == \"google\" or extra == \"extra-proxy\") and python_version < \"3.14\"", proxy-dev = "python_version >= \"3.10\" and python_version < \"3.14\""}
[package.dependencies]
google-auth = ">=2.14.1,<3.0.0"
@@ -1909,7 +1906,7 @@ files = [
{file = "google_auth-2.43.0-py2.py3-none-any.whl", hash = "sha256:af628ba6fa493f75c7e9dbe9373d148ca9f4399b5ea29976519e0a3848eddd16"},
{file = "google_auth-2.43.0.tar.gz", hash = "sha256:88228eee5fc21b62a1b5fe773ca15e67778cb07dc8363adcb4a8827b52d81483"},
]
markers = {main = "python_version >= \"3.10\" and (extra == \"extra-proxy\" or extra == \"google\" or extra == \"mlflow\") or extra == \"google\" or extra == \"extra-proxy\"", proxy-dev = "python_version >= \"3.10\""}
markers = {main = "extra == \"google\" or extra == \"extra-proxy\" or python_version >= \"3.10\"", proxy-dev = "python_version >= \"3.10\""}
[package.dependencies]
cachetools = ">=2.0.0,<7.0"
@@ -2081,11 +2078,11 @@ files = [
]
[package.dependencies]
google-api-core = {version = ">=1.34.1,<2.0.dev0 || >=2.11.dev0,<3.0.0.dev0", extras = ["grpc"]}
google-auth = ">=2.14.1,<2.24.0 || >2.24.0,<2.25.0 || >2.25.0,<3.0.0.dev0"
grpc-google-iam-v1 = ">=0.12.4,<1.0.0.dev0"
proto-plus = ">=1.22.3,<2.0.0.dev0"
protobuf = ">=3.20.2,<4.21.0 || >4.21.0,<4.21.1 || >4.21.1,<4.21.2 || >4.21.2,<4.21.3 || >4.21.3,<4.21.4 || >4.21.4,<4.21.5 || >4.21.5,<6.0.0.dev0"
google-api-core = {version = ">=1.34.1,<2.0.dev0 || >=2.11.dev0,<3.0.0dev", extras = ["grpc"]}
google-auth = ">=2.14.1,<2.24.0 || >2.24.0,<2.25.0 || >2.25.0,<3.0.0dev"
grpc-google-iam-v1 = ">=0.12.4,<1.0.0dev"
proto-plus = ">=1.22.3,<2.0.0dev"
protobuf = ">=3.20.2,<4.21.0 || >4.21.0,<4.21.1 || >4.21.1,<4.21.2 || >4.21.2,<4.21.3 || >4.21.3,<4.21.4 || >4.21.4,<4.21.5 || >4.21.5,<6.0.0dev"
[[package]]
name = "google-cloud-resource-manager"
@@ -2267,7 +2264,7 @@ files = [
{file = "googleapis_common_protos-1.72.0-py3-none-any.whl", hash = "sha256:4299c5a82d5ae1a9702ada957347726b167f9f8d1fc352477702a1e851ff4038"},
{file = "googleapis_common_protos-1.72.0.tar.gz", hash = "sha256:e55a601c1b32b52d7a3e65f43563e2aa61bcd737998ee672ac9b951cd49319f5"},
]
markers = {main = "python_version >= \"3.10\" and (extra == \"extra-proxy\" or extra == \"google\") or extra == \"google\" or extra == \"extra-proxy\""}
markers = {main = "extra == \"google\" or extra == \"extra-proxy\" or python_version >= \"3.10\""}
[package.dependencies]
grpcio = {version = ">=1.44.0,<2.0.0", optional = true, markers = "extra == \"grpc\""}
@@ -2676,11 +2673,11 @@ description = "Consume Server-Sent Event (SSE) messages with HTTPX."
optional = false
python-versions = ">=3.9"
groups = ["main", "proxy-dev"]
markers = "python_version >= \"3.10\""
files = [
{file = "httpx_sse-0.4.3-py3-none-any.whl", hash = "sha256:0ac1c9fe3c0afad2e0ebb25a934a59f4c7823b60792691f779fad2c5568830fc"},
{file = "httpx_sse-0.4.3.tar.gz", hash = "sha256:9b1ed0127459a66014aec3c56bebd93da3c1bc8bb6618c8082039a44889a755d"},
]
markers = {main = "python_version >= \"3.10\" and (extra == \"proxy\" or extra == \"extra-proxy\")", proxy-dev = "python_version >= \"3.10\""}
[[package]]
name = "huey"
@@ -3045,7 +3042,7 @@ files = [
[package.dependencies]
attrs = ">=22.2.0"
jsonschema-specifications = ">=2023.3.6"
jsonschema-specifications = ">=2023.03.6"
referencing = ">=0.28.4"
rpds-py = ">=0.7.1"
@@ -3716,7 +3713,6 @@ files = [
{file = "msal-1.34.0-py3-none-any.whl", hash = "sha256:f669b1644e4950115da7a176441b0e13ec2975c29528d8b9e81316023676d6e1"},
{file = "msal-1.34.0.tar.gz", hash = "sha256:76ba83b716ea5a6d75b0279c0ac353a0e05b820ca1f6682c0eb7f45190c43c2f"},
]
markers = {main = "extra == \"proxy\" or extra == \"extra-proxy\""}
[package.dependencies]
cryptography = ">=2.5,<49"
@@ -3737,7 +3733,6 @@ files = [
{file = "msal_extensions-1.3.1-py3-none-any.whl", hash = "sha256:96d3de4d034504e969ac5e85bae8106c8373b5c6568e4c8fa7af2eca9dbe6bca"},
{file = "msal_extensions-1.3.1.tar.gz", hash = "sha256:c5b0fd10f65ef62b5f1d62f4251d51cbcaf003fcedae8c91b040a488614be1a4"},
]
markers = {main = "extra == \"proxy\" or extra == \"extra-proxy\""}
[package.dependencies]
msal = ">=1.29,<2"
@@ -3988,7 +3983,6 @@ files = [
{file = "nodeenv-1.9.1-py2.py3-none-any.whl", hash = "sha256:ba11c9782d29c27c70ffbdda2d7415098754709be8a7056d79a737cd901155c9"},
{file = "nodeenv-1.9.1.tar.gz", hash = "sha256:6ec12890a2dab7946721edbfbcd91f3319c6ccc9aec47be7c7e6b7011ee6645f"},
]
markers = {main = "extra == \"extra-proxy\""}
[[package]]
name = "numpy"
@@ -4111,7 +4105,7 @@ files = [
{file = "opentelemetry_api-1.39.1-py3-none-any.whl", hash = "sha256:2edd8463432a7f8443edce90972169b195e7d6a05500cd29e6d13898187c9950"},
{file = "opentelemetry_api-1.39.1.tar.gz", hash = "sha256:fbde8c80e1b937a2c61f20347e91c0c18a1940cecf012d62e65a7caf08967c9c"},
]
markers = {main = "python_version >= \"3.10\" and extra == \"mlflow\""}
markers = {main = "python_version >= \"3.10\""}
[package.dependencies]
importlib-metadata = ">=6.0,<8.8.0"
@@ -4226,7 +4220,7 @@ files = [
{file = "opentelemetry_sdk-1.39.1-py3-none-any.whl", hash = "sha256:4d5482c478513ecb0a5d938dcc61394e647066e0cc2676bee9f3af3f3f45f01c"},
{file = "opentelemetry_sdk-1.39.1.tar.gz", hash = "sha256:cf4d4563caf7bff906c9f7967e2be22d0d6b349b908be0d90fb21c8e9c995cc6"},
]
markers = {main = "python_version >= \"3.10\" and extra == \"mlflow\""}
markers = {main = "python_version >= \"3.10\""}
[package.dependencies]
opentelemetry-api = "1.39.1"
@@ -4244,7 +4238,7 @@ files = [
{file = "opentelemetry_semantic_conventions-0.60b1-py3-none-any.whl", hash = "sha256:9fa8c8b0c110da289809292b0591220d3a7b53c1526a23021e977d68597893fb"},
{file = "opentelemetry_semantic_conventions-0.60b1.tar.gz", hash = "sha256:87c228b5a0669b748c76d76df6c364c369c28f1c465e50f661e39737e84bc953"},
]
markers = {main = "python_version >= \"3.10\" and extra == \"mlflow\""}
markers = {main = "python_version >= \"3.10\""}
[package.dependencies]
opentelemetry-api = "1.39.1"
@@ -4728,7 +4722,6 @@ files = [
{file = "prisma-0.11.0-py3-none-any.whl", hash = "sha256:22bb869e59a2968b99f3483bb417717273ffbc569fd1e9ceed95e5614cbaf53a"},
{file = "prisma-0.11.0.tar.gz", hash = "sha256:3f2f2fd2361e1ec5ff655f2a04c7860c2f2a5bc4c91f78ca9c5c6349735bf693"},
]
markers = {main = "extra == \"extra-proxy\""}
[package.dependencies]
click = ">=7.1.2"
@@ -4902,7 +4895,7 @@ files = [
{file = "proto_plus-1.26.1-py3-none-any.whl", hash = "sha256:13285478c2dcf2abb829db158e1047e2f1e8d63a077d94263c2b88b043c75a66"},
{file = "proto_plus-1.26.1.tar.gz", hash = "sha256:21a515a4c4c0088a773899e23c7bbade3d18f9c66c73edd4c7ee3816bc96a012"},
]
markers = {main = "extra == \"google\" or extra == \"extra-proxy\"", proxy-dev = "python_version >= \"3.10\""}
markers = {main = "extra == \"google\" or extra == \"extra-proxy\" or python_version >= \"3.10\"", proxy-dev = "python_version >= \"3.10\""}
[package.dependencies]
protobuf = ">=3.19.0,<7.0.0"
@@ -4930,7 +4923,7 @@ files = [
{file = "protobuf-5.29.5-py3-none-any.whl", hash = "sha256:6cf42630262c59b2d8de33954443d94b746c952b01434fc58a417fdbd2e84bd5"},
{file = "protobuf-5.29.5.tar.gz", hash = "sha256:bc1463bafd4b0929216c35f437a8e28731a2b7fe3d98bb77a600efced5a15c84"},
]
markers = {main = "python_version >= \"3.10\" and (extra == \"extra-proxy\" or extra == \"google\" or extra == \"mlflow\") or extra == \"google\" or extra == \"extra-proxy\""}
markers = {main = "extra == \"google\" or extra == \"extra-proxy\" or python_version >= \"3.10\""}
[[package]]
name = "psutil"
@@ -5090,7 +5083,7 @@ files = [
{file = "pyasn1-0.6.1-py3-none-any.whl", hash = "sha256:0d632f46f2ba09143da3a8afe9e33fb6f92fa2320ab7e886e2d0f7672af84629"},
{file = "pyasn1-0.6.1.tar.gz", hash = "sha256:6f580d2bdd84365380830acf45550f2511469f673cb4a5ae3857a3170128b034"},
]
markers = {main = "python_version >= \"3.10\" and (extra == \"extra-proxy\" or extra == \"google\" or extra == \"mlflow\") or extra == \"google\" or extra == \"extra-proxy\"", proxy-dev = "python_version >= \"3.10\""}
markers = {main = "extra == \"google\" or extra == \"extra-proxy\" or python_version >= \"3.10\"", proxy-dev = "python_version >= \"3.10\""}
[[package]]
name = "pyasn1-modules"
@@ -5103,7 +5096,7 @@ files = [
{file = "pyasn1_modules-0.4.2-py3-none-any.whl", hash = "sha256:29253a9207ce32b64c3ac6600edc75368f98473906e8fd1043bd6b5b1de2c14a"},
{file = "pyasn1_modules-0.4.2.tar.gz", hash = "sha256:677091de870a80aae844b1ca6134f54652fa2c8c5a52aa396440ac3106e941e6"},
]
markers = {main = "python_version >= \"3.10\" and (extra == \"extra-proxy\" or extra == \"google\" or extra == \"mlflow\") or extra == \"google\" or extra == \"extra-proxy\"", proxy-dev = "python_version >= \"3.10\""}
markers = {main = "extra == \"google\" or extra == \"extra-proxy\" or python_version >= \"3.10\"", proxy-dev = "python_version >= \"3.10\""}
[package.dependencies]
pyasn1 = ">=0.6.1,<0.7.0"
@@ -5131,7 +5124,7 @@ files = [
{file = "pycparser-2.23-py3-none-any.whl", hash = "sha256:e5c6e8d3fbad53479cab09ac03729e0a9faf2bee3db8208a550daf5af81a5934"},
{file = "pycparser-2.23.tar.gz", hash = "sha256:78816d4f24add8f10a06d6f05b4d424ad9e96cfebf68a4ddc99c65c0720d00c2"},
]
markers = {main = "implementation_name != \"PyPy\" and (platform_python_implementation != \"PyPy\" or extra == \"proxy\") and (python_version >= \"3.10\" or extra == \"proxy\" or extra == \"extra-proxy\") and (extra == \"proxy\" or extra == \"extra-proxy\" or extra == \"mlflow\")", dev = "platform_python_implementation != \"PyPy\" and implementation_name != \"PyPy\"", proxy-dev = "platform_python_implementation != \"PyPy\" and implementation_name != \"PyPy\""}
markers = {main = "implementation_name != \"PyPy\" and (platform_python_implementation != \"PyPy\" or extra == \"proxy\")", dev = "platform_python_implementation != \"PyPy\" and implementation_name != \"PyPy\"", proxy-dev = "platform_python_implementation != \"PyPy\" and implementation_name != \"PyPy\""}
[[package]]
name = "pydantic"
@@ -5354,7 +5347,6 @@ files = [
{file = "PyJWT-2.10.1-py3-none-any.whl", hash = "sha256:dcdd193e30abefd5debf142f9adfcdd2b58004e644f25406ffaebd50bd98dacb"},
{file = "pyjwt-2.10.1.tar.gz", hash = "sha256:3cc5772eb20009233caf06e9d8a0577824723b44e6648ee0a2aedb6cf9381953"},
]
markers = {main = "(python_version <= \"3.13\" or extra == \"proxy\" or extra == \"extra-proxy\") and (extra == \"extra-proxy\" or extra == \"proxy\")"}
[package.dependencies]
cryptography = {version = ">=3.4.0", optional = true, markers = "extra == \"crypto\""}
@@ -6284,7 +6276,7 @@ files = [
{file = "rsa-4.9.1-py3-none-any.whl", hash = "sha256:68635866661c6836b8d39430f97a996acbd61bfa49406748ea243539fe239762"},
{file = "rsa-4.9.1.tar.gz", hash = "sha256:e7bdbfdb5497da4c07dfd35530e1a902659db6ff241e39d9953cad06ebd0ae75"},
]
markers = {main = "python_version >= \"3.10\" and (extra == \"extra-proxy\" or extra == \"google\" or extra == \"mlflow\") or extra == \"google\" or extra == \"extra-proxy\"", proxy-dev = "python_version >= \"3.10\""}
markers = {main = "extra == \"google\" or extra == \"extra-proxy\" or python_version >= \"3.10\"", proxy-dev = "python_version >= \"3.10\""}
[package.dependencies]
pyasn1 = ">=0.1.3"
@@ -6330,10 +6322,10 @@ files = [
]
[package.dependencies]
botocore = ">=1.37.4,<2.0a0"
botocore = ">=1.37.4,<2.0a.0"
[package.extras]
crt = ["botocore[crt] (>=1.37.4,<2.0a0)"]
crt = ["botocore[crt] (>=1.37.4,<2.0a.0)"]
[[package]]
name = "scikit-learn"
@@ -6486,9 +6478,9 @@ tornado = ">=6.4.2,<7"
urllib3 = ">=1.26,<3"
[package.extras]
all = ["boto3 (>=1.34.98,<2)", "botocore (>=1.34.110,<2)", "cohere (>=5.9.4,<6.0)", "dagger-io (>=0.1.1) ; python_version >= \"3.11\"", "fastembed (>=0.3.0,<0.4) ; python_version < \"3.13\"", "google-cloud-aiplatform (>=1.45.0,<2)", "ipykernel (>=6.25.0,<7)", "llama-cpp-python (>=0.2.28,<0.2.86) ; python_version < \"3.13\"", "mistralai (>=0.0.12,<0.1.0)", "mypy (>=1.7.1,<2)", "ollama (>=0.1.7)", "pillow (>=10.2.0,<11.0.0) ; python_version < \"3.13\"", "pinecone[asyncio] (>=7.0.0,<8.0.0)", "psycopg[binary] (>=3.1.0,<4)", "pytest (>=8.2,<9.0)", "pytest-asyncio (>=0.24.0,<0.25)", "pytest-cov (>=4.1.0,<5)", "pytest-mock (>=3.12.0,<4)", "pytest-timeout", "pytest-xdist (>=3.5.0,<4)", "python-dotenv (>=1.0.0,<2)", "qdrant-client (>=1.11.1,<2)", "requests-mock (>=1.12.1,<2)", "ruff (>=0.11.2,<0.12)", "sentence-transformers (>=5.0.0) ; python_version < \"3.13\"", "tokenizers (>=0.19) ; python_version < \"3.13\"", "torch (>=2.6.0) ; python_version < \"3.13\"", "torchvision (>=0.17.0) ; python_version < \"3.13\"", "transformers (>=4.36.2) ; python_version < \"3.13\"", "types-pyyaml (>=6.0.12.12,<7)", "types-requests (>=2.31.0,<3)"]
all = ["boto3 (>=1.34.98,<2)", "botocore (>=1.34.110,<2)", "cohere (>=5.9.4,<6.00)", "dagger-io (>=0.1.1) ; python_version >= \"3.11\"", "fastembed (>=0.3.0,<0.4) ; python_version < \"3.13\"", "google-cloud-aiplatform (>=1.45.0,<2)", "ipykernel (>=6.25.0,<7)", "llama-cpp-python (>=0.2.28,<0.2.86) ; python_version < \"3.13\"", "mistralai (>=0.0.12,<0.1.0)", "mypy (>=1.7.1,<2)", "ollama (>=0.1.7)", "pillow (>=10.2.0,<11.0.0) ; python_version < \"3.13\"", "pinecone[asyncio] (>=7.0.0,<8.0.0)", "psycopg[binary] (>=3.1.0,<4)", "pytest (>=8.2,<9.0)", "pytest-asyncio (>=0.24.0,<0.25)", "pytest-cov (>=4.1.0,<5)", "pytest-mock (>=3.12.0,<4)", "pytest-timeout", "pytest-xdist (>=3.5.0,<4)", "python-dotenv (>=1.0.0,<2)", "qdrant-client (>=1.11.1,<2)", "requests-mock (>=1.12.1,<2)", "ruff (>=0.11.2,<0.12)", "sentence-transformers (>=5.0.0) ; python_version < \"3.13\"", "tokenizers (>=0.19) ; python_version < \"3.13\"", "torch (>=2.6.0) ; python_version < \"3.13\"", "torchvision (>=0.17.0) ; python_version < \"3.13\"", "transformers (>=4.36.2) ; python_version < \"3.13\"", "types-pyyaml (>=6.0.12.12,<7)", "types-requests (>=2.31.0,<3)"]
bedrock = ["boto3 (>=1.34.98,<2)", "botocore (>=1.34.110,<2)"]
cohere = ["cohere (>=5.9.4,<6.0)"]
cohere = ["cohere (>=5.9.4,<6.00)"]
dev = ["dagger-io (>=0.1.1) ; python_version >= \"3.11\"", "ipykernel (>=6.25.0,<7)", "mypy (>=1.7.1,<2)", "pytest (>=8.2,<9.0)", "pytest-asyncio (>=0.24.0,<0.25)", "pytest-cov (>=4.1.0,<5)", "pytest-mock (>=3.12.0,<4)", "pytest-timeout", "pytest-xdist (>=3.5.0,<4)", "python-dotenv (>=1.0.0,<2)", "requests-mock (>=1.12.1,<2)", "ruff (>=0.11.2,<0.12)", "types-pyyaml (>=6.0.12.12,<7)", "types-requests (>=2.31.0,<3)"]
docs = ["pydoc-markdown (>=4.8.2) ; python_version < \"3.12\""]
fastembed = ["fastembed (>=0.3.0,<0.4) ; python_version < \"3.13\""]
@@ -7216,7 +7208,6 @@ files = [
{file = "tomlkit-0.13.3-py3-none-any.whl", hash = "sha256:c89c649d79ee40629a9fda55f8ace8c6a1b42deb912b2a8fd8d942ddadb606b0"},
{file = "tomlkit-0.13.3.tar.gz", hash = "sha256:430cf247ee57df2b94ee3fbe588e71d362a941ebb545dec29b53961d61add2a1"},
]
markers = {main = "extra == \"extra-proxy\""}
[[package]]
name = "tornado"
@@ -7496,15 +7487,15 @@ zstd = ["zstandard (>=0.18.0)"]
[[package]]
name = "uvicorn"
version = "0.31.1"
version = "0.39.0"
description = "The lightning-fast ASGI server."
optional = true
python-versions = ">=3.8"
python-versions = ">=3.9"
groups = ["main"]
markers = "python_version >= \"3.10\" and (extra == \"mlflow\" or extra == \"proxy\") or extra == \"proxy\""
markers = "python_version == \"3.9\" and extra == \"proxy\""
files = [
{file = "uvicorn-0.31.1-py3-none-any.whl", hash = "sha256:adc42d9cac80cf3e51af97c1851648066841e7cfb6993a4ca8de29ac1548ed41"},
{file = "uvicorn-0.31.1.tar.gz", hash = "sha256:f5167919867b161b7bcaf32646c6a94cdbd4c3aa2eb5c17d36bb9aa5cfd8c493"},
{file = "uvicorn-0.39.0-py3-none-any.whl", hash = "sha256:7beec21bd2693562b386285b188a7963b06853c0d006302b3e4cfed950c9929a"},
{file = "uvicorn-0.39.0.tar.gz", hash = "sha256:610512b19baa93423d2892d7823741f6d27717b642c8964000d7194dded19302"},
]
[package.dependencies]
@@ -7513,7 +7504,28 @@ h11 = ">=0.8"
typing-extensions = {version = ">=4.0", markers = "python_version < \"3.11\""}
[package.extras]
standard = ["colorama (>=0.4) ; sys_platform == \"win32\"", "httptools (>=0.5.0)", "python-dotenv (>=0.13)", "pyyaml (>=5.1)", "uvloop (>=0.14.0,!=0.15.0,!=0.15.1) ; sys_platform != \"win32\" and sys_platform != \"cygwin\" and platform_python_implementation != \"PyPy\"", "watchfiles (>=0.13)", "websockets (>=10.4)"]
standard = ["colorama (>=0.4) ; sys_platform == \"win32\"", "httptools (>=0.6.3)", "python-dotenv (>=0.13)", "pyyaml (>=5.1)", "uvloop (>=0.15.1) ; sys_platform != \"win32\" and sys_platform != \"cygwin\" and platform_python_implementation != \"PyPy\"", "watchfiles (>=0.13)", "websockets (>=10.4)"]
[[package]]
name = "uvicorn"
version = "0.41.0"
description = "The lightning-fast ASGI server."
optional = true
python-versions = ">=3.10"
groups = ["main"]
markers = "python_version >= \"3.10\" and (extra == \"mlflow\" or extra == \"proxy\")"
files = [
{file = "uvicorn-0.41.0-py3-none-any.whl", hash = "sha256:29e35b1d2c36a04b9e180d4007ede3bcb32a85fbdfd6c6aeb3f26839de088187"},
{file = "uvicorn-0.41.0.tar.gz", hash = "sha256:09d11cf7008da33113824ee5a1c6422d89fbc2ff476540d69a34c87fab8b571a"},
]
[package.dependencies]
click = ">=7.0"
h11 = ">=0.8"
typing-extensions = {version = ">=4.0", markers = "python_version < \"3.11\""}
[package.extras]
standard = ["colorama (>=0.4) ; sys_platform == \"win32\"", "httptools (>=0.6.3)", "python-dotenv (>=0.13)", "pyyaml (>=5.1)", "uvloop (>=0.15.1) ; sys_platform != \"win32\" and sys_platform != \"cygwin\" and platform_python_implementation != \"PyPy\"", "watchfiles (>=0.20)", "websockets (>=10.4)"]
[[package]]
name = "uvloop"
@@ -7968,4 +7980,4 @@ utils = ["numpydoc"]
[metadata]
lock-version = "2.1"
python-versions = ">=3.9,<4.0"
content-hash = "3dd495ee4e23d7cb750525c4f364ee96a4ef34fa9d9d5c4ed07b5432c0925d48"
content-hash = "97936ece74659668c195c2c05ff36c6edd24c5c18b52ab3a5bb8a3cd2f329e5d"
+2 -2
View File
@@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm"
version = "1.81.13"
version = "1.81.14"
description = "Library to easily interface with LLM API providers"
authors = ["BerriAI"]
license = "MIT"
@@ -183,7 +183,7 @@ requires = ["poetry-core", "wheel"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
version = "1.81.13"
version = "1.81.14"
version_files = [
"pyproject.toml:^version"
]
+21 -5
View File
@@ -35,8 +35,23 @@ EXCLUDED_TERMINAL_VARS = {
"ALACRITTY_SOCKET",
}
# Directories to skip (dependencies, venvs, caches) - only scan litellm source
SKIP_DIRS = {
".venv",
"venv",
"__pycache__",
".git",
"node_modules",
"site-packages",
".eggs",
"dist",
"build",
}
# Walk through all files in the litellm repo to find references of os.getenv() and litellm.get_secret()
for root, dirs, files in os.walk(repo_base):
# Skip dependency/venv directories - prevents picking up env vars from installed packages
dirs[:] = [d for d in dirs if d not in SKIP_DIRS]
for file in files:
if file.endswith(".py"): # Only process Python files
file_path = os.path.join(root, file)
@@ -83,12 +98,13 @@ try:
)
print(f"general_settings_section: {general_settings_section}")
if general_settings_section:
# Extract the table rows, which contain the documented keys
# Extract the table rows - only first column (key name) from each row
table_content = general_settings_section.group(1)
doc_key_pattern = re.compile(
r"\|\s*([^\|]+?)\s*\|"
) # Capture the key from each row of the table
documented_keys.update(doc_key_pattern.findall(table_content))
for line in table_content.split("\n"):
# Match | KEY_NAME | description | - capture first column only
match = re.match(r"^\|\s*([A-Z_][A-Z0-9_]*)\s*\|", line)
if match:
documented_keys.add(match.group(1).strip())
except Exception as e:
raise Exception(
f"Error reading documentation: {e}, \n repo base - {os.listdir(repo_base)}"
@@ -0,0 +1,81 @@
"""Unit tests for LiteLLMAgentModelResolver - litellm_agent/ prefix model resolution."""
from unittest.mock import MagicMock
import pytest
from litellm.integrations.litellm_agent import LiteLLMAgentModelResolver
class TestLiteLLMAgentModelResolver:
def test_get_chat_completion_prompt_strips_prefix(self):
"""Verify get_chat_completion_prompt strips litellm_agent/ prefix from model."""
resolver = LiteLLMAgentModelResolver()
messages = [{"role": "user", "content": "Hello"}]
resolved_model, out_messages, out_params = resolver.get_chat_completion_prompt(
model="litellm_agent/gpt-3.5-turbo",
messages=messages,
non_default_params={},
prompt_id=None,
prompt_variables=None,
dynamic_callback_params={},
)
assert resolved_model == "gpt-3.5-turbo"
assert out_messages == messages
assert out_params == {}
def test_get_chat_completion_prompt_preserves_rest_of_model(self):
"""Verify model name after prefix is preserved (e.g. openai/gpt-3.5-turbo)."""
resolver = LiteLLMAgentModelResolver()
messages = [{"role": "user", "content": "Test"}]
resolved_model, _, _ = resolver.get_chat_completion_prompt(
model="litellm_agent/openai/gpt-3.5-turbo",
messages=messages,
non_default_params={},
prompt_id=None,
prompt_variables=None,
dynamic_callback_params={},
)
assert resolved_model == "openai/gpt-3.5-turbo"
def test_get_chat_completion_prompt_respects_ignore_prompt_manager_model(self):
"""Verify model is unchanged when ignore_prompt_manager_model is True."""
resolver = LiteLLMAgentModelResolver()
messages = [{"role": "user", "content": "Hello"}]
resolved_model, _, _ = resolver.get_chat_completion_prompt(
model="litellm_agent/gpt-3.5-turbo",
messages=messages,
non_default_params={},
prompt_id=None,
prompt_variables=None,
dynamic_callback_params={},
ignore_prompt_manager_model=True,
)
assert resolved_model == "litellm_agent/gpt-3.5-turbo"
@pytest.mark.asyncio
async def test_async_get_chat_completion_prompt_strips_prefix(self):
"""Verify async_get_chat_completion_prompt strips prefix."""
resolver = LiteLLMAgentModelResolver()
messages = [{"role": "user", "content": "Hello"}]
resolved_model, out_messages, _ = (
await resolver.async_get_chat_completion_prompt(
model="litellm_agent/gpt-3.5-turbo",
messages=messages,
non_default_params={},
prompt_id=None,
prompt_variables=None,
dynamic_callback_params={},
litellm_logging_obj=MagicMock(),
)
)
assert resolved_model == "gpt-3.5-turbo"
assert out_messages == messages
@@ -10,6 +10,7 @@ sys.path.insert(
) # Adds the parent directory to the system path
from litellm.litellm_core_utils.prompt_templates.common_utils import (
add_system_prompt_to_messages,
get_format_from_file_id,
handle_any_messages_to_chat_completion_str_messages_conversion,
split_concatenated_json_objects,
@@ -128,6 +129,60 @@ def test_handle_any_messages_to_chat_completion_str_messages_conversion_complex(
assert result[0]["input"] == json.dumps(message)
def test_add_system_prompt_to_messages_prepend():
"""Adds system prompt at beginning when no system message exists."""
messages = [
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Hi there"},
]
result = add_system_prompt_to_messages(messages, "You are a helpful assistant.")
assert result == [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Hi there"},
]
def test_add_system_prompt_to_messages_empty_prompt_unchanged():
"""Returns messages unchanged when system_prompt is empty."""
messages = [{"role": "user", "content": "Hello"}]
assert add_system_prompt_to_messages(messages, "") == messages
assert add_system_prompt_to_messages(messages, None) == messages
def test_add_system_prompt_to_messages_merge_with_first_system():
"""Merges new prompt into first system message when merge_with_first_system=True."""
messages = [
{"role": "system", "content": "Existing system prompt."},
{"role": "user", "content": "Hello"},
]
result = add_system_prompt_to_messages(
messages, "You are helpful.", merge_with_first_system=True
)
assert result == [
{"role": "system", "content": "You are helpful.\n\nExisting system prompt."},
{"role": "user", "content": "Hello"},
]
def test_add_system_prompt_to_messages_merge_with_first_system_adds_new_when_no_system():
"""When merge_with_first_system=True but no system message, adds new one at start."""
messages = [{"role": "user", "content": "Hello"}]
result = add_system_prompt_to_messages(
messages, "You are helpful.", merge_with_first_system=True
)
assert result == [
{"role": "system", "content": "You are helpful."},
{"role": "user", "content": "Hello"},
]
def test_add_system_prompt_to_messages_empty_list():
"""Adds system prompt to empty messages list."""
result = add_system_prompt_to_messages([], "You are helpful.")
assert result == [{"role": "system", "content": "You are helpful."}]
def test_convert_prefix_message_to_non_prefix_messages():
from litellm.litellm_core_utils.prompt_templates.common_utils import (
convert_prefix_message_to_non_prefix_messages,
@@ -1122,6 +1122,44 @@ def test_get_final_response_obj_with_empty_response_obj_and_list_init():
assert result[1].name == "Object2"
def test_get_usage_as_dict():
"""
Test get_usage_as_dict returns usage as plain dict from response_obj or combined_usage_object.
"""
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.types.utils import Usage
# Test case 1: None response_obj returns empty usage dict
result = StandardLoggingPayloadSetup.get_usage_as_dict(response_obj=None)
assert result == {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
# Test case 2: Empty response_obj returns empty usage dict
result = StandardLoggingPayloadSetup.get_usage_as_dict(response_obj={})
assert result == {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
# Test case 3: combined_usage_object takes priority
combined = Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15)
result = StandardLoggingPayloadSetup.get_usage_as_dict(
response_obj={"usage": {"prompt_tokens": 1, "completion_tokens": 1}},
combined_usage_object=combined,
)
assert result["prompt_tokens"] == 10
assert result["completion_tokens"] == 5
assert result["total_tokens"] == 15
# Test case 4: response_obj with usage dict
result = StandardLoggingPayloadSetup.get_usage_as_dict(
response_obj={"usage": {"prompt_tokens": 20, "completion_tokens": 30}}
)
assert result == {"prompt_tokens": 20, "completion_tokens": 30}
# Test case 5: response_obj with no usage key returns empty
result = StandardLoggingPayloadSetup.get_usage_as_dict(
response_obj={"id": "resp-1", "choices": []}
)
assert result == {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
def test_append_system_prompt_messages():
"""
Test append_system_prompt_messages prepends system message from kwargs to messages list.
@@ -2924,3 +2924,23 @@ def test_fast_mode_parameter_mapping():
assert "speed" in result
assert result["speed"] == "fast"
def test_map_openai_params_max_tokens_normalized_to_int():
"""
Test that map_openai_params normalizes max_tokens to an integer (e.g. 0.7 -> 1).
"""
config = AnthropicConfig()
non_default_params = {"max_tokens": 0.7}
optional_params = {}
result = config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model="claude-3-5-sonnet-20241022",
drop_params=False,
)
assert "max_tokens" in result
assert result["max_tokens"] == 1
@@ -0,0 +1,78 @@
"""Tests for the response-rejection custom guardrail code (input_type response, block on refusal)."""
import pytest
from fastapi import HTTPException
from litellm.proxy.guardrails.guardrail_hooks.custom_code import (
RESPONSE_REJECTION_GUARDRAIL_CODE, CustomCodeGuardrail)
@pytest.fixture
def response_rejection_guardrail():
"""Guardrail instance using the response-rejection custom code."""
return CustomCodeGuardrail(
guardrail_name="response_rejection",
custom_code=RESPONSE_REJECTION_GUARDRAIL_CODE,
)
@pytest.mark.asyncio
async def test_response_rejection_allows_request_input_type(response_rejection_guardrail):
"""Should allow when input_type is 'request' (no response check)."""
result = await response_rejection_guardrail.apply_guardrail(
inputs={"texts": ["some user message"]},
request_data={},
input_type="request",
)
assert result == {"texts": ["some user message"]}
@pytest.mark.asyncio
async def test_response_rejection_allows_helpful_response(response_rejection_guardrail):
"""Should allow when response text does not contain rejection phrases."""
result = await response_rejection_guardrail.apply_guardrail(
inputs={"texts": ["Here is how you can do that: step 1, step 2."]},
request_data={},
input_type="response",
)
assert result["texts"] == ["Here is how you can do that: step 1, step 2."]
@pytest.mark.asyncio
async def test_response_rejection_blocks_refusal_phrase(response_rejection_guardrail):
"""Should block when response contains a known rejection phrase."""
with pytest.raises(HTTPException) as exc_info:
await response_rejection_guardrail.apply_guardrail(
inputs={"texts": ["That's not something I can help with."]},
request_data={},
input_type="response",
)
assert exc_info.value.status_code == 400
detail = exc_info.value.detail
assert isinstance(detail, dict)
assert "error" in detail
assert "rejected" in detail["error"].lower() or "reject" in detail["error"].lower()
assert detail.get("guardrail") == "response_rejection"
assert detail.get("detection_info", {}).get("matched_phrase") is not None
@pytest.mark.asyncio
async def test_response_rejection_blocks_case_insensitive(response_rejection_guardrail):
"""Should block on refusal phrase regardless of case."""
with pytest.raises(HTTPException):
await response_rejection_guardrail.apply_guardrail(
inputs={"texts": ["I'M SORRY, I CAN'T do that."]},
request_data={},
input_type="response",
)
@pytest.mark.asyncio
async def test_response_rejection_empty_texts_allowed(response_rejection_guardrail):
"""Should allow when texts is empty or missing."""
result = await response_rejection_guardrail.apply_guardrail(
inputs={},
request_data={},
input_type="response",
)
assert result == {}
@@ -1,6 +1,7 @@
"use client";
import { useState, useEffect } from "react";
import AgentBuilderView from "@/components/playground/chat_ui/AgentBuilderView";
import ChatUI from "@/components/playground/chat_ui/ChatUI";
import CompareUI from "@/components/playground/compareUI/CompareUI";
import ComplianceUI from "@/components/playground/complianceUI/ComplianceUI";
@@ -39,6 +40,7 @@ export default function PlaygroundPage() {
<Tab>Chat</Tab>
<Tab>Compare</Tab>
<Tab>Compliance</Tab>
<Tab>Agent Builder (Experimental)</Tab>
</TabList>
<TabPanels className="h-full">
<TabPanel className="h-full">
@@ -57,6 +59,17 @@ export default function PlaygroundPage() {
<TabPanel className="h-full">
<ComplianceUI accessToken={accessToken} disabledPersonalKeyCreation={disabledPersonalKeyCreation} />
</TabPanel>
<TabPanel className="h-full">
<AgentBuilderView
accessToken={accessToken}
token={token}
userID={userId}
userRole={userRole}
disabledPersonalKeyCreation={disabledPersonalKeyCreation}
proxySettings={proxySettings}
customProxyBaseUrl={proxySettings?.LITELLM_UI_API_DOC_BASE_URL ?? proxySettings?.PROXY_BASE_URL}
/>
</TabPanel>
</TabPanels>
</TabGroup>
);
@@ -5453,6 +5453,8 @@ export interface TestPoliciesAndGuardrailsRequest {
inputs_list?: GuardrailInputs[] | null;
request_data?: Record<string, unknown>;
input_type?: "request" | "response";
/** When set, backend runs chat completion with this model/agent per input and includes agent_response in each result. */
agent_id?: string | null;
}
export interface GuardrailErrorEntry {
@@ -5460,16 +5462,24 @@ export interface GuardrailErrorEntry {
message: string;
}
export interface TestPoliciesAndGuardrailsResultItem {
inputs: Record<string, unknown>;
guardrail_errors: GuardrailErrorEntry[];
/** Present when request included agent_id; serialized chat completion response. */
agent_response?: Record<string, unknown>;
}
export interface TestPoliciesAndGuardrailsResponse {
inputs?: Record<string, unknown>;
guardrail_errors?: GuardrailErrorEntry[];
/** Present when inputs_list was used; one result per input. */
results?: Array<{ inputs: Record<string, unknown>; guardrail_errors: GuardrailErrorEntry[] }>;
results?: TestPoliciesAndGuardrailsResultItem[];
}
export const testPoliciesAndGuardrails = async (
accessToken: string,
body: TestPoliciesAndGuardrailsRequest
body: TestPoliciesAndGuardrailsRequest,
signal?: AbortSignal
): Promise<TestPoliciesAndGuardrailsResponse> => {
try {
const url = proxyBaseUrl
@@ -5477,6 +5487,7 @@ export const testPoliciesAndGuardrails = async (
: `/utils/test_policies_and_guardrails`;
const response = await fetch(url, {
method: "POST",
signal,
headers: {
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
"Content-Type": "application/json",
@@ -5488,6 +5499,7 @@ export const testPoliciesAndGuardrails = async (
inputs_list: body.inputs_list ?? null,
request_data: body.request_data ?? {},
input_type: body.input_type ?? "request",
agent_id: body.agent_id ?? null,
}),
});
@@ -0,0 +1,712 @@
"use client";
import { CommentOutlined, DeleteOutlined, ExperimentOutlined, LinkOutlined, PlusOutlined, RobotOutlined, SaveOutlined } from "@ant-design/icons";
import { Button, Input, Modal, Select, Spin, Tabs } from "antd";
import React, { useCallback, useEffect, useState } from "react";
import CodeBlock from "@/app/(dashboard)/api-reference/components/CodeBlock";
import NotificationsManager from "../../molecules/notifications_manager";
import { keyCreateCall, modelCreateCall, modelDeleteCall, modelPatchUpdateCall, proxyBaseUrl } from "../../networking";
import { fetchMCPServers } from "../../networking";
import { MCPServer } from "../../mcp_tools/types";
import { AgentModel, fetchAvailableAgentModels, MCPToolEntry } from "../llm_calls/fetch_agents";
import { fetchAvailableModels, ModelGroup } from "../llm_calls/fetch_models";
import ComplianceUI from "../complianceUI/ComplianceUI";
import ChatUI from "./ChatUI";
const { TextArea } = Input;
export interface AgentBuilderViewProps {
accessToken: string | null;
token: string | null;
userID: string | null;
userRole: string | null;
disabledPersonalKeyCreation?: boolean;
proxySettings?: {
PROXY_BASE_URL?: string;
LITELLM_UI_API_DOC_BASE_URL?: string | null;
};
apiKey?: string;
customProxyBaseUrl?: string;
}
const NEW_AGENT_ID = "__new__";
function getConnectTabBaseUrl(
proxySettings: AgentBuilderViewProps["proxySettings"],
customProxyBaseUrl?: string,
): string {
const customDocBaseUrl = proxySettings?.LITELLM_UI_API_DOC_BASE_URL;
if (customDocBaseUrl && customDocBaseUrl.trim()) return customDocBaseUrl;
if (proxySettings?.PROXY_BASE_URL) return proxySettings.PROXY_BASE_URL;
if (customProxyBaseUrl?.trim()) return customProxyBaseUrl;
return "<your_proxy_base_url>";
}
interface ConnectTabContentProps {
agentName: string;
proxySettings: AgentBuilderViewProps["proxySettings"];
customProxyBaseUrl?: string;
accessToken: string | null;
userID: string | null;
disabledPersonalKeyCreation: boolean;
creatingKey: boolean;
createdKeyValue: string | null;
onCreateKey: () => void;
}
function ConnectTabContent({
agentName,
proxySettings,
customProxyBaseUrl,
disabledPersonalKeyCreation,
creatingKey,
createdKeyValue,
onCreateKey,
}: ConnectTabContentProps) {
const baseUrl = proxyBaseUrl ?? getConnectTabBaseUrl(proxySettings, customProxyBaseUrl);
const apiKeyForCurl =
createdKeyValue ?
createdKeyValue.startsWith("Bearer ") ? createdKeyValue : `Bearer ${createdKeyValue}`
: "Bearer sk-1234";
const curlExample = `curl -L -X POST '${baseUrl}/v1/chat/completions' \\
-H 'x-litellm-api-key: ${apiKeyForCurl}' \\
-d '{
"model": "${agentName}",
"stream": true,
"stream_options": {
"include_usage": true
},
"messages": [
{
"role": "user",
"content": "hey"
}
]
}'`;
return (
<div className="mx-auto max-w-3xl space-y-6">
<div>
<h3 className="text-sm font-semibold text-gray-900 mb-1">Proxy base URL</h3>
<p className="text-sm text-gray-600 font-mono bg-gray-50 px-2 py-1.5 rounded border border-gray-200 break-all">
{baseUrl}
</p>
</div>
<div>
<h3 className="text-sm font-semibold text-gray-900 mb-2">Call your agent (cURL)</h3>
<CodeBlock code={curlExample} language="bash" />
</div>
<div className="rounded-lg border border-gray-200 bg-gray-50 p-4">
<h3 className="text-sm font-semibold text-gray-900 mb-2">Create a key for this agent</h3>
<p className="text-sm text-gray-600 mb-3">
Create a virtual key that can only call this agent. The key will be scoped to you (user_id) and restricted to
the model <span className="font-mono text-gray-800">{agentName}</span>.
</p>
<Button
type="primary"
onClick={onCreateKey}
loading={creatingKey}
disabled={disabledPersonalKeyCreation}
>
Create key for this agent
</Button>
{disabledPersonalKeyCreation && (
<p className="text-xs text-amber-600 mt-2">Key creation is disabled for your account.</p>
)}
{createdKeyValue && (
<p className="text-xs text-green-700 mt-2">
Key created. It is shown in the cURL example above copy the snippet to use it.
</p>
)}
</div>
</div>
);
}
function getAgentModelId(agent: AgentModel): string | null {
const info = agent.model_info as { id?: string } | null | undefined;
return info?.id ?? null;
}
function parseUnderlyingModel(litellmModel: string | undefined): string | undefined {
if (!litellmModel || !litellmModel.startsWith("litellm_agent/")) return undefined;
return litellmModel.slice("litellm_agent/".length) || undefined;
}
const MCP_TOOLS_PREFIX = "litellm_proxy/mcp/";
function buildToolsFromServerIds(serverIds: string[], servers: MCPServer[]): MCPToolEntry[] {
return serverIds.map((serverId) => {
const server = servers.find((s) => s.server_id === serverId);
const serverName = server?.alias || server?.server_name || serverId;
return {
type: "mcp",
server_label: "litellm",
server_url: `${MCP_TOOLS_PREFIX}${serverName}`,
require_approval: "never",
};
});
}
function getServerIdsFromTools(tools: MCPToolEntry[], servers: MCPServer[]): string[] {
return tools
.filter((t) => t.type === "mcp" && t.server_url?.startsWith(MCP_TOOLS_PREFIX))
.map((t) => {
const suffix = t.server_url.slice(MCP_TOOLS_PREFIX.length);
const server = servers.find((s) => (s.alias || s.server_name || s.server_id) === suffix);
return server?.server_id;
})
.filter((id): id is string => id != null);
}
export default function AgentBuilderView({
accessToken,
token,
userID,
userRole,
disabledPersonalKeyCreation = false,
proxySettings,
apiKey,
customProxyBaseUrl,
}: AgentBuilderViewProps) {
const [agentModels, setAgentModels] = useState<AgentModel[]>([]);
const [modelGroups, setModelGroups] = useState<ModelGroup[]>([]);
const [loadingAgents, setLoadingAgents] = useState(true);
const [selectedId, setSelectedId] = useState<string | null>(null);
const [activeTab, setActiveTab] = useState<"configure" | "chat" | "test" | "connect">("configure");
const [creatingKey, setCreatingKey] = useState(false);
const [createdKeyValue, setCreatedKeyValue] = useState<string | null>(null);
// Draft for new agent
const [draftName, setDraftName] = useState("");
const [draftSystemPrompt, setDraftSystemPrompt] = useState("");
const [draftUnderlyingModel, setDraftUnderlyingModel] = useState<string | undefined>(undefined);
const [draftTemperature, setDraftTemperature] = useState(0.7);
const [draftMaxTokens, setDraftMaxTokens] = useState(4096);
const [draftTools, setDraftTools] = useState<MCPToolEntry[]>([]);
const [mcpServers, setMCPServers] = useState<MCPServer[]>([]);
const [loadingMCPServers, setLoadingMCPServers] = useState(false);
const [saving, setSaving] = useState(false);
const [deleting, setDeleting] = useState(false);
const effectiveApiKey = apiKey || accessToken || "";
const selectedAgent = selectedId === NEW_AGENT_ID ? null : agentModels.find((a) => a.model_name === selectedId) ?? null;
const isNewAgent = selectedId === NEW_AGENT_ID;
const selectedAgentModelId = selectedAgent ? getAgentModelId(selectedAgent) : null;
const loadAgents = useCallback(async () => {
if (!accessToken || !userID || !userRole) return;
setLoadingAgents(true);
try {
const list = await fetchAvailableAgentModels(accessToken, userID, userRole);
setAgentModels(list);
if (!selectedId || (selectedId !== NEW_AGENT_ID && !list.some((a) => a.model_name === selectedId))) {
setSelectedId(list.length > 0 ? list[0].model_name : null);
}
} catch (e) {
console.error(e);
NotificationsManager.fromBackend("Failed to load agents");
} finally {
setLoadingAgents(false);
}
}, [accessToken, userID, userRole]);
const loadModels = useCallback(async () => {
if (!effectiveApiKey) return;
try {
const models = await fetchAvailableModels(effectiveApiKey);
setModelGroups(models);
if (!draftUnderlyingModel && models.length > 0) {
setDraftUnderlyingModel(models[0].model_group);
}
} catch (e) {
console.error(e);
}
}, [effectiveApiKey]);
useEffect(() => {
loadAgents();
}, [loadAgents]);
useEffect(() => {
loadModels();
}, [loadModels]);
const loadMCPServers = useCallback(async () => {
if (!effectiveApiKey) return;
setLoadingMCPServers(true);
try {
const servers = await fetchMCPServers(effectiveApiKey);
setMCPServers(Array.isArray(servers) ? servers : (servers as { data?: MCPServer[] })?.data ?? []);
} catch (e) {
console.error("Error fetching MCP servers:", e);
} finally {
setLoadingMCPServers(false);
}
}, [effectiveApiKey]);
useEffect(() => {
loadMCPServers();
}, [loadMCPServers]);
// Clear created key when switching to another agent
useEffect(() => {
setCreatedKeyValue(null);
}, [selectedId]);
// Sync draft fields when selecting an existing agent
useEffect(() => {
if (selectedAgent && !isNewAgent) {
setDraftName(selectedAgent.model_name);
setDraftSystemPrompt(selectedAgent.litellm_params?.litellm_system_prompt ?? "");
const underlying = parseUnderlyingModel(selectedAgent.litellm_params?.model);
setDraftUnderlyingModel(underlying ?? modelGroups[0]?.model_group);
const p = selectedAgent.litellm_params as { temperature?: number; max_tokens?: number } | undefined;
setDraftTemperature(typeof p?.temperature === "number" ? p.temperature : 0.7);
setDraftMaxTokens(typeof p?.max_tokens === "number" ? p.max_tokens : 4096);
const rawTools = selectedAgent.litellm_params?.tools;
const tools: MCPToolEntry[] = Array.isArray(rawTools)
? rawTools.filter((t): t is MCPToolEntry => t && typeof t === "object" && (t as MCPToolEntry).type === "mcp" && typeof (t as MCPToolEntry).server_url === "string")
: [];
setDraftTools(tools);
}
}, [selectedId, isNewAgent, selectedAgent?.model_name, selectedAgent?.litellm_params?.tools]);
const selectedMCPServerIds = getServerIdsFromTools(draftTools, mcpServers);
const handleMCPServerChange = (serverIds: string[]) => {
setDraftTools(buildToolsFromServerIds(serverIds, mcpServers));
};
const handleAddAgent = () => {
setSelectedId(NEW_AGENT_ID);
setDraftName("");
setDraftSystemPrompt("You are a helpful assistant.");
setDraftUnderlyingModel(modelGroups[0]?.model_group);
setDraftTemperature(0.7);
setDraftMaxTokens(4096);
setDraftTools([]);
setActiveTab("configure");
};
const handleSaveAgent = async () => {
if (!accessToken || !draftName?.trim() || !draftUnderlyingModel) {
NotificationsManager.fromBackend("Name and underlying model are required");
return;
}
setSaving(true);
try {
await modelCreateCall(accessToken, {
model_name: draftName.trim(),
litellm_params: {
model: `litellm_agent/${draftUnderlyingModel}`,
litellm_system_prompt: draftSystemPrompt.trim() || undefined,
temperature: draftTemperature,
max_tokens: draftMaxTokens,
tools: draftTools,
},
model_info: {},
});
const newName = draftName.trim();
await loadAgents();
setSelectedId(newName);
setActiveTab("chat");
} catch (e) {
NotificationsManager.fromBackend("Failed to save agent");
} finally {
setSaving(false);
}
};
const handleUpdateAgent = async () => {
if (!accessToken || !selectedAgent || !selectedAgentModelId || !draftName?.trim() || !draftUnderlyingModel) {
NotificationsManager.fromBackend("Name and underlying model are required");
return;
}
setSaving(true);
try {
await modelPatchUpdateCall(
accessToken,
{
model_name: draftName.trim(),
litellm_params: {
model: `litellm_agent/${draftUnderlyingModel}`,
litellm_system_prompt: draftSystemPrompt.trim() || undefined,
temperature: draftTemperature,
max_tokens: draftMaxTokens,
tools: draftTools,
},
model_info: selectedAgent.model_info ?? {},
},
selectedAgentModelId,
);
NotificationsManager.success("Agent updated successfully");
await loadAgents();
setSelectedId(draftName.trim());
} catch (e) {
NotificationsManager.fromBackend("Failed to update agent");
} finally {
setSaving(false);
}
};
const handleCreateKeyForAgent = async () => {
if (!accessToken || !userID || !selectedAgent) return;
setCreatingKey(true);
setCreatedKeyValue(null);
try {
const response = await keyCreateCall(accessToken, userID, {
models: [selectedAgent.model_name],
key_alias: `Agent: ${selectedAgent.model_name}`,
});
const keyValue = response?.key ?? null;
if (keyValue) {
setCreatedKeyValue(keyValue);
NotificationsManager.success("Virtual key created. Use it in the curl example below.");
} else {
NotificationsManager.fromBackend("Key created but value not returned");
}
} catch (e) {
NotificationsManager.fromBackend("Failed to create key for agent");
} finally {
setCreatingKey(false);
}
};
const handleDeleteAgent = () => {
if (!selectedAgent || !selectedAgentModelId || !accessToken) return;
Modal.confirm({
title: "Delete agent",
content: `Are you sure you want to delete "${selectedAgent.model_name}"? This cannot be undone.`,
okText: "Delete",
okType: "danger",
cancelText: "Cancel",
onOk: async () => {
setDeleting(true);
try {
await modelDeleteCall(accessToken, selectedAgentModelId);
NotificationsManager.success("Agent deleted");
await loadAgents();
const remaining = agentModels.filter((a) => a.model_name !== selectedAgent.model_name);
setSelectedId(remaining.length > 0 ? remaining[0].model_name : null);
} catch (e) {
NotificationsManager.fromBackend("Failed to delete agent");
} finally {
setDeleting(false);
}
},
});
};
if (!accessToken || !userID || !userRole) {
return (
<div className="flex h-full items-center justify-center p-8 text-gray-500">
Sign in to use Agent Builder.
</div>
);
}
return (
<div className="flex h-full flex-col bg-white text-gray-900">
<div className="flex flex-shrink-0 flex-col border-b border-gray-200">
<div className="flex h-12 items-center justify-between px-4">
<span className="text-sm font-medium text-gray-900">Agent Builder</span>
{isNewAgent ? (
<Button
type="primary"
icon={<SaveOutlined />}
onClick={handleSaveAgent}
loading={saving}
disabled={!draftName?.trim() || !draftUnderlyingModel}
>
Save Agent
</Button>
) : (
<span className="text-xs text-gray-500">Build Agents that pass your compliance requirements.</span>
)}
</div>
<div className="flex items-center gap-2 border-t border-amber-200 bg-amber-50 px-4 py-2 text-xs text-amber-800">
<ExperimentOutlined className="flex-shrink-0 text-amber-600" />
<span>
Agent Builder is experimental and may change or be removed without notice. Wed love your feedbackemail us at{" "}
<a href="mailto:product@berri.ai" className="font-medium text-amber-900 underline hover:text-amber-700">
product@berri.ai
</a>
.
</span>
</div>
</div>
<div className="flex flex-1 overflow-hidden">
{/* Roster */}
<div className="w-60 flex-shrink-0 border-r border-gray-200 bg-white flex flex-col">
<div className="flex items-center justify-between border-b border-gray-200 p-3">
<span className="text-xs font-semibold uppercase tracking-wide text-gray-500">Agents</span>
<Button type="text" size="small" icon={<PlusOutlined />} onClick={handleAddAgent} aria-label="Add agent" />
</div>
<div className="flex-1 overflow-y-auto p-2">
{loadingAgents ? (
<div className="flex justify-center py-4">
<Spin size="small" />
</div>
) : (
<>
{agentModels.map((agent) => (
<button
key={agent.model_name}
type="button"
onClick={() => setSelectedId(agent.model_name)}
className={`mb-1 w-full rounded-md border-l-2 px-3 py-2 text-left text-sm transition-colors ${
selectedId === agent.model_name
? "border-blue-500 bg-blue-50 text-blue-800"
: "border-transparent hover:bg-gray-50"
}`}
>
<div className="font-medium truncate">{agent.model_name}</div>
<div className="text-[10px] text-gray-500 truncate">litellm_agent</div>
</button>
))}
<button
type="button"
onClick={handleAddAgent}
className="mb-1 w-full rounded-md border border-dashed border-gray-300 px-3 py-2 text-left text-sm text-gray-500 hover:border-blue-400 hover:bg-blue-50/50 hover:text-gray-700"
>
<PlusOutlined className="mr-1" /> New agent
</button>
</>
)}
</div>
</div>
{/* Main content */}
<div className="flex flex-1 flex-col overflow-hidden">
{selectedId === null && !isNewAgent && agentModels.length === 0 && !loadingAgents && (
<div className="flex flex-1 items-center justify-center p-8 text-gray-500">
No agents yet. Add an agent to get started.
</div>
)}
{(selectedId !== null || isNewAgent) && (
<>
<Tabs
activeKey={activeTab}
onChange={(k) => setActiveTab(k as "configure" | "chat" | "test" | "connect")}
className="flex-1 overflow-hidden [&_.ant-tabs-content]:h-full [&_.ant-tabs-tabpane]:h-full [&_.ant-tabs-nav]:pl-4"
items={[
{
key: "configure",
label: (
<span>
<RobotOutlined className="mr-1" /> Configure
</span>
),
children: (
<div className="h-full overflow-y-auto p-6">
{(isNewAgent || selectedAgent) ? (
<div className="mx-auto max-w-xl space-y-4">
{!selectedAgentModelId && selectedAgent && (
<div className="rounded border border-amber-200 bg-amber-50 px-3 py-2 text-xs text-amber-800">
This agent cannot be updated or deleted here (missing model id). Manage it from Models &amp; Endpoints.
</div>
)}
<div>
<label className="mb-1 block text-sm font-medium text-gray-700">Agent name</label>
<Input
value={draftName}
onChange={(e) => setDraftName(e.target.value)}
placeholder="My Agent"
/>
</div>
<div>
<label className="mb-1 block text-sm font-medium text-gray-700">System prompt</label>
<TextArea
value={draftSystemPrompt}
onChange={(e) => setDraftSystemPrompt(e.target.value)}
placeholder="You are a helpful assistant..."
rows={6}
/>
</div>
<div>
<label className="mb-1 block text-sm font-medium text-gray-700">Underlying LLM</label>
<Select
value={draftUnderlyingModel}
onChange={setDraftUnderlyingModel}
className="w-full"
options={modelGroups.map((m) => ({ value: m.model_group, label: m.model_group }))}
placeholder="Select model"
/>
</div>
<div className="grid grid-cols-2 gap-4">
<div>
<label className="mb-1 block text-sm font-medium text-gray-700">Temperature</label>
<Input
type="number"
min={0}
max={2}
step={0.1}
value={draftTemperature}
onChange={(e) => setDraftTemperature(Number(e.target.value))}
/>
</div>
<div>
<label className="mb-1 block text-sm font-medium text-gray-700">Max tokens</label>
<Input
type="number"
min={1}
value={draftMaxTokens}
onChange={(e) => setDraftMaxTokens(Number(e.target.value))}
/>
</div>
</div>
<div>
<label className="mb-1 block text-sm font-medium text-gray-700">MCP servers</label>
<Select
mode="multiple"
placeholder="Select MCP servers to attach (same format as chat completions API)"
value={selectedMCPServerIds}
onChange={handleMCPServerChange}
loading={loadingMCPServers}
className="w-full"
allowClear
showSearch
optionFilterProp="label"
options={mcpServers.map((s) => ({
value: s.server_id,
label: s.alias || s.server_name || s.server_id,
}))}
/>
{selectedAgent && draftTools.length > 0 && (
<p className="mt-1 text-xs text-gray-500">
{draftTools.length} MCP server{draftTools.length !== 1 ? "s" : ""} saved. Use the same <code className="rounded bg-gray-100 px-1">tools</code> array in chat completions when calling this agent.
</p>
)}
</div>
{selectedAgent && (
<div className="flex flex-wrap items-center gap-2 pt-2">
{selectedAgentModelId && (
<>
<Button
type="primary"
icon={<SaveOutlined />}
onClick={handleUpdateAgent}
loading={saving}
disabled={!draftName?.trim() || !draftUnderlyingModel}
>
Update Agent
</Button>
<Button
type="default"
danger
icon={<DeleteOutlined />}
onClick={handleDeleteAgent}
loading={deleting}
>
Delete
</Button>
</>
)}
<Button type="primary" icon={<CommentOutlined />} onClick={() => setActiveTab("chat")}>
Test in Chat
</Button>
</div>
)}
</div>
) : null}
</div>
),
},
{
key: "chat",
label: (
<span>
<CommentOutlined className="mr-1" /> Chat
</span>
),
disabled: isNewAgent,
children: (
<div className="flex h-full flex-col min-h-0">
{selectedAgent ? (
<ChatUI
key={selectedAgent.model_name}
simplified
fixedModel={selectedAgent.model_name}
accessToken={accessToken}
token={token}
userRole={userRole}
userID={userID}
disabledPersonalKeyCreation={disabledPersonalKeyCreation}
proxySettings={proxySettings}
/>
) : (
<div className="flex flex-1 items-center justify-center text-gray-500">
Save an agent first to test in Chat.
</div>
)}
</div>
),
},
{
key: "test",
label: (
<span>
<ExperimentOutlined className="mr-1" /> Batch Test
</span>
),
disabled: isNewAgent,
children: (
<div className="flex h-full flex-col min-h-0">
{selectedAgent ? (
<ComplianceUI
accessToken={accessToken}
disabledPersonalKeyCreation={disabledPersonalKeyCreation}
backendMode="chat_completions"
fixedModel={selectedAgent.model_name}
proxySettings={proxySettings}
/>
) : (
<div className="flex flex-1 items-center justify-center text-gray-500">
Select an agent to run batch tests.
</div>
)}
</div>
),
},
{
key: "connect",
label: (
<span>
<LinkOutlined className="mr-1" /> Connect
</span>
),
disabled: isNewAgent,
children: (
<div className="h-full overflow-y-auto p-6">
{selectedAgent ? (
<ConnectTabContent
agentName={selectedAgent.model_name}
proxySettings={proxySettings}
customProxyBaseUrl={customProxyBaseUrl}
accessToken={accessToken}
userID={userID}
disabledPersonalKeyCreation={disabledPersonalKeyCreation}
creatingKey={creatingKey}
createdKeyValue={createdKeyValue}
onCreateKey={handleCreateKeyForAgent}
/>
) : (
<div className="flex flex-1 items-center justify-center text-gray-500">
Select an agent to see how to connect.
</div>
)}
</div>
),
},
]}
/>
</>
)}
</div>
</div>
</div>
);
}
@@ -84,6 +84,10 @@ interface ChatUIProps {
PROXY_BASE_URL?: string;
LITELLM_UI_API_DOC_BASE_URL?: string | null;
};
/** When true, hide configuration sidebar and use fixedModel only (e.g. embedded in Agent Builder). */
simplified?: boolean;
/** When simplified is true, use this as the model and do not show model selector. */
fixedModel?: string;
}
const MCP_SUPPORTED_ENDPOINTS = new Set<EndpointType>([
@@ -99,6 +103,8 @@ const ChatUI: React.FC<ChatUIProps> = ({
userID,
disabledPersonalKeyCreation,
proxySettings,
simplified = false,
fixedModel,
}) => {
const [mcpServers, setMCPServers] = useState<MCPServer[]>([]);
const [selectedMCPServers, setSelectedMCPServers] = useState<string[]>(() => {
@@ -140,6 +146,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
);
const [inputMessage, setInputMessage] = useState("");
const [chatHistory, setChatHistory] = useState<MessageType[]>(() => {
if (simplified) return [];
try {
const saved = sessionStorage.getItem("chatHistory");
return saved ? JSON.parse(saved) : [];
@@ -148,7 +155,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
return [];
}
});
const [selectedModel, setSelectedModel] = useState<string | undefined>(undefined);
const [selectedModel, setSelectedModel] = useState<string | undefined>(simplified ? fixedModel : undefined);
const [showCustomModelInput, setShowCustomModelInput] = useState<boolean>(false);
const [modelInfo, setModelInfo] = useState<ModelGroup[]>([]);
const [agentInfo, setAgentInfo] = useState<Agent[]>([]);
@@ -252,6 +259,14 @@ const ChatUI: React.FC<ChatUIProps> = ({
}
};
// When simplified, keep selectedModel and endpointType in sync with fixedModel / chat-only
useEffect(() => {
if (simplified && fixedModel) {
setSelectedModel(fixedModel);
setEndpointType(EndpointType.CHAT);
}
}, [simplified, fixedModel]);
// Fetch tools for a specific server
const loadServerTools = async (serverId: string) => {
const userApiKey = apiKeySource === "session" ? accessToken : apiKey;
@@ -312,6 +327,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
]);
useEffect(() => {
if (simplified) return; // Do not persist chat history in simplified (embedded) mode
const handler = setTimeout(() => {
sessionStorage.setItem("chatHistory", JSON.stringify(chatHistory));
}, 500); // Debounce by 500ms
@@ -319,7 +335,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
return () => {
clearTimeout(handler);
};
}, [chatHistory]);
}, [chatHistory, simplified]);
useEffect(() => {
sessionStorage.setItem("apiKeySource", JSON.stringify(apiKeySource));
@@ -334,10 +350,12 @@ const ChatUI: React.FC<ChatUIProps> = ({
sessionStorage.setItem("selectedVoice", selectedVoice);
sessionStorage.removeItem("selectedMCPTools"); // Clean up old key
if (selectedModel) {
sessionStorage.setItem("selectedModel", selectedModel);
} else {
sessionStorage.removeItem("selectedModel");
if (!simplified) {
if (selectedModel) {
sessionStorage.setItem("selectedModel", selectedModel);
} else {
sessionStorage.removeItem("selectedModel");
}
}
if (messageTraceId) {
sessionStorage.setItem("messageTraceId", messageTraceId);
@@ -352,6 +370,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
sessionStorage.setItem("useApiSessionManagement", JSON.stringify(useApiSessionManagement));
// Note: codeInterpreterEnabled and selectedContainerId are persisted by useCodeInterpreter hook
}, [
simplified,
apiKeySource,
apiKey,
selectedModel,
@@ -375,7 +394,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
return;
}
// Fetch model info and set the default selected model
// Fetch model info and set the default selected model (skip in simplified mode; we use fixedModel)
const loadModels = async () => {
try {
if (!userApiKey) {
@@ -400,9 +419,11 @@ const ChatUI: React.FC<ChatUIProps> = ({
}
};
loadModels();
if (!simplified) {
loadModels();
}
loadMCPServers();
}, [accessToken, userID, userRole, apiKeySource, apiKey, token]);
}, [accessToken, userID, userRole, apiKeySource, apiKey, token, simplified]);
// Load tools when MCP direct mode has a server selected
useEffect(() => {
@@ -868,7 +889,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
return;
}
const effectiveApiKey = apiKeySource === "session" ? accessToken : apiKey;
const effectiveApiKey = simplified ? accessToken : apiKeySource === "session" ? accessToken : apiKey;
if (!effectiveApiKey) {
NotificationsManager.fromBackend("Please provide a Virtual Key or select Current UI Session");
@@ -959,6 +980,10 @@ const ChatUI: React.FC<ChatUIProps> = ({
newUserMessage,
];
const requestProxyBaseUrl =
simplified && proxySettings
? (proxySettings.LITELLM_UI_API_DOC_BASE_URL ?? proxySettings.PROXY_BASE_URL ?? undefined)
: (customProxyBaseUrl || undefined);
await makeOpenAIChatCompletionRequest(
apiChatHistory,
(chunk, model) => updateTextUI("assistant", chunk, model),
@@ -979,7 +1004,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
useAdvancedParams ? temperature : undefined,
useAdvancedParams ? maxTokens : undefined,
updateTotalLatency,
customProxyBaseUrl || undefined,
requestProxyBaseUrl,
mcpServers,
mcpServerToolRestrictions,
handleMCPEvent,
@@ -1207,9 +1232,11 @@ const ChatUI: React.FC<ChatUIProps> = ({
handleRemoveResponsesImage(); // Clear any uploaded images for responses
handleRemoveChatImage(); // Clear any uploaded images for chat completions
handleRemoveAudio(); // Clear any uploaded audio for transcription
sessionStorage.removeItem("chatHistory");
sessionStorage.removeItem("messageTraceId");
sessionStorage.removeItem("responsesSessionId");
if (!simplified) {
sessionStorage.removeItem("chatHistory");
sessionStorage.removeItem("messageTraceId");
sessionStorage.removeItem("responsesSessionId");
}
NotificationsManager.success("Chat history cleared.");
};
@@ -1246,10 +1273,11 @@ const ChatUI: React.FC<ChatUIProps> = ({
const antIcon = <LoadingOutlined style={{ fontSize: 24 }} spin />;
return (
<div className="w-full p-4 pb-0 bg-white">
<Card className="w-full rounded-xl shadow-md overflow-hidden">
<div className="flex h-[80vh] w-full gap-4">
{/* Left Sidebar with Controls */}
<div className={`w-full bg-white ${simplified ? "h-full flex flex-col" : "p-4 pb-0"}`}>
<Card className={`w-full rounded-xl shadow-md overflow-hidden ${simplified ? "h-full flex flex-col" : ""}`}>
<div className={`flex w-full gap-4 ${simplified ? "h-full" : "h-[80vh]"}`}>
{/* Left Sidebar with Controls - hidden in simplified mode */}
{!simplified && (
<div className="w-1/4 p-4 bg-gray-50 overflow-y-auto">
<Title className="text-xl font-semibold mb-6 mt-2">Configurations</Title>
<div className="space-y-4">
@@ -1794,11 +1822,12 @@ const ChatUI: React.FC<ChatUIProps> = ({
)}
</div>
</div>
)}
{/* Main Chat Area */}
<div className="w-3/4 flex flex-col bg-white">
<div className={`flex flex-col bg-white ${simplified ? "flex-1 w-full" : "w-3/4"}`}>
<div className="p-4 border-b border-gray-200 flex justify-between items-center">
<Title className="text-xl font-semibold mb-0">Test Key</Title>
<Title className="text-xl font-semibold mb-0">{simplified ? "Chat" : "Test Key"}</Title>
<div className="flex gap-2">
<TremorButton
onClick={clearChatHistory}
@@ -1807,6 +1836,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
>
Clear Chat
</TremorButton>
{!simplified && (
<TremorButton
onClick={() => setIsGetCodeModalVisible(true)}
className="bg-gray-100 hover:bg-gray-200 text-gray-700 border-gray-300"
@@ -1814,6 +1844,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
>
Get Code
</TremorButton>
)}
</div>
</div>
<div className="flex-1 overflow-auto p-4 pb-0">
@@ -11,6 +11,7 @@ import {
getPoliciesList,
testPoliciesAndGuardrails,
} from "@/components/networking";
import { makeOpenAIChatCompletionRequest } from "../llm_calls/chat_completion";
import {
AlertTriangle,
BarChart3,
@@ -38,6 +39,7 @@ import {
Send,
Shield,
Smile,
Square,
Trash2,
TrendingDown,
Upload,
@@ -110,11 +112,23 @@ interface GuardrailOption {
interface ComplianceUIProps {
accessToken: string | null;
disabledPersonalKeyCreation?: boolean;
/** When "chat_completions", use /chat/completions with fixedModel instead of test_policies_and_guardrails. */
backendMode?: "policies" | "chat_completions";
/** Required when backendMode is "chat_completions"; model name for chat completions (e.g. selected agent). */
fixedModel?: string;
/** Used when backendMode is "chat_completions" for the request base URL. */
proxySettings?: {
PROXY_BASE_URL?: string;
LITELLM_UI_API_DOC_BASE_URL?: string | null;
};
}
export default function ComplianceUI({
accessToken,
disabledPersonalKeyCreation,
backendMode = "policies",
fixedModel,
proxySettings,
}: ComplianceUIProps) {
const frameworks = getFrameworks();
@@ -148,6 +162,7 @@ export default function ComplianceUI({
const [isRunning, setIsRunning] = useState(false);
const [resultFilter, setResultFilter] = useState<ResultFilter>("all");
const [expandedResults, setExpandedResults] = useState<Set<string>>(new Set());
const batchAbortControllerRef = useRef<AbortController | null>(null);
useEffect(() => {
if (!accessToken) return;
@@ -432,6 +447,9 @@ export default function ComplianceUI({
if (csvInputRef.current) csvInputRef.current.value = "";
};
const requestProxyBaseUrl =
proxySettings?.LITELLM_UI_API_DOC_BASE_URL ?? proxySettings?.PROXY_BASE_URL ?? undefined;
const runQuickTest = useCallback(async () => {
if (!quickTestInput.trim() || !accessToken) return;
const text = quickTestInput.trim();
@@ -445,44 +463,82 @@ export default function ComplianceUI({
setQuickTestInput("");
setIsQuickTesting(true);
try {
const { inputs, guardrail_errors = [] } = await testPoliciesAndGuardrails(
accessToken,
{
policy_names:
selectedPolicies.length > 0 ? selectedPolicies : undefined,
guardrail_names:
selectedGuardrails.length > 0 ? selectedGuardrails : undefined,
inputs: { texts: [text] },
request_data: {},
input_type: "request",
}
);
const result: "blocked" | "allowed" =
guardrail_errors.length > 0 ? "blocked" : "allowed";
const triggeredBy =
guardrail_errors.length > 0
? guardrail_errors
.map((e) => `${e.guardrail_name}: ${e.message}`)
.join("; ")
: undefined;
const returnedText =
Array.isArray(inputs?.texts) && inputs.texts.length > 0
? inputs.texts[0]
: undefined;
const displayText =
result === "blocked"
? `Blocked — ${triggeredBy ?? "content filter"}`
: "Allowed — no policy or guardrail violations detected.";
const sysMsg: QuickTestMessage = {
id: `msg-${Date.now()}-sys`,
type: "system",
text: displayText,
result,
triggeredBy,
returnedText,
timestamp: new Date(),
};
setQuickTestMessages((prev) => [...prev, sysMsg]);
if (backendMode === "chat_completions" && fixedModel) {
let fullResponse = "";
await makeOpenAIChatCompletionRequest(
[{ role: "user", content: text }],
(chunk: string) => {
fullResponse += chunk;
},
fixedModel,
accessToken,
undefined,
undefined,
undefined,
undefined,
undefined,
undefined,
undefined, // vector_store_ids (param 11)
selectedGuardrails.length > 0 ? selectedGuardrails : undefined,
selectedPolicies.length > 0 ? selectedPolicies : undefined,
undefined,
undefined,
undefined,
undefined,
undefined,
undefined,
requestProxyBaseUrl,
undefined
);
const sysMsg: QuickTestMessage = {
id: `msg-${Date.now()}-sys`,
type: "system",
text: "Allowed — model response received.",
result: "allowed",
returnedText: fullResponse,
timestamp: new Date(),
};
setQuickTestMessages((prev) => [...prev, sysMsg]);
} else {
const { inputs, guardrail_errors = [] } = await testPoliciesAndGuardrails(
accessToken,
{
policy_names:
selectedPolicies.length > 0 ? selectedPolicies : undefined,
guardrail_names:
selectedGuardrails.length > 0 ? selectedGuardrails : undefined,
inputs: { texts: [text] },
request_data: {},
input_type: "request",
}
);
const result: "blocked" | "allowed" =
guardrail_errors.length > 0 ? "blocked" : "allowed";
const triggeredBy =
guardrail_errors.length > 0
? guardrail_errors
.map((e) => `${e.guardrail_name}: ${e.message}`)
.join("; ")
: undefined;
const returnedText =
Array.isArray(inputs?.texts) && inputs.texts.length > 0
? inputs.texts[0]
: undefined;
const displayText =
result === "blocked"
? `Blocked — ${triggeredBy ?? "content filter"}`
: "Allowed — no policy or guardrail violations detected.";
const sysMsg: QuickTestMessage = {
id: `msg-${Date.now()}-sys`,
type: "system",
text: displayText,
result,
triggeredBy,
returnedText,
timestamp: new Date(),
};
setQuickTestMessages((prev) => [...prev, sysMsg]);
}
} catch (err) {
const errorMessage = err instanceof Error ? err.message : String(err);
const sysMsg: QuickTestMessage = {
@@ -502,6 +558,9 @@ export default function ComplianceUI({
quickTestInput,
selectedPolicies,
selectedGuardrails,
backendMode,
fixedModel,
requestProxyBaseUrl,
]);
const handleQuickTestKeyDown = (e: React.KeyboardEvent<HTMLTextAreaElement>) => {
@@ -513,6 +572,9 @@ export default function ComplianceUI({
const runTests = useCallback(async () => {
if (selectedPromptIds.size === 0 || !accessToken) return;
const controller = new AbortController();
batchAbortControllerRef.current = controller;
const signal = controller.signal;
setIsRunning(true);
setResultFilter("all");
setRightTab("batch-results");
@@ -533,16 +595,21 @@ export default function ComplianceUI({
}));
setTestResults(pendingResults);
try {
const inputsList = allTexts.map((text) => ({ texts: [text] }));
const response = await testPoliciesAndGuardrails(accessToken, {
policy_names:
selectedPolicies.length > 0 ? selectedPolicies : undefined,
guardrail_names:
selectedGuardrails.length > 0 ? selectedGuardrails : undefined,
inputs_list: inputsList,
request_data: {},
input_type: "request",
});
const useAgentId = backendMode === "chat_completions" && fixedModel;
const response = await testPoliciesAndGuardrails(
accessToken,
{
policy_names:
selectedPolicies.length > 0 ? selectedPolicies : undefined,
guardrail_names:
selectedGuardrails.length > 0 ? selectedGuardrails : undefined,
inputs_list: allTexts.map((text) => ({ texts: [text] })),
request_data: {},
input_type: "request",
...(useAgentId ? { agent_id: fixedModel } : {}),
},
signal
);
const results = response.results ?? [];
setTestResults(
pendingResults.map((row, index) => {
@@ -556,10 +623,17 @@ export default function ComplianceUI({
.map((e) => `${e.guardrail_name}: ${e.message}`)
.join("; ")
: undefined;
const returnedText =
Array.isArray(item?.inputs?.texts) && item.inputs.texts.length > 0
? item.inputs.texts[0]
: undefined;
let returnedText: string | undefined;
if (item?.agent_response != null) {
const choices = (item.agent_response as { choices?: Array<{ message?: { content?: string } }> }).choices;
returnedText =
Array.isArray(choices) && choices[0]?.message?.content != null
? String(choices[0].message.content)
: undefined;
}
if (returnedText === undefined && Array.isArray(item?.inputs?.texts) && item.inputs.texts.length > 0) {
returnedText = item.inputs.texts[0] as string;
}
return {
...row,
actualResult,
@@ -573,6 +647,10 @@ export default function ComplianceUI({
})
);
} catch (err) {
if (err instanceof Error && err.name === "AbortError") {
// Stopped by user; leave partial results as-is (already set in loop)
return;
}
const errorMessage = err instanceof Error ? err.message : String(err);
setTestResults(
pendingResults.map((row) => ({
@@ -583,14 +661,19 @@ export default function ComplianceUI({
status: "complete" as const,
}))
);
} finally {
setIsRunning(false);
batchAbortControllerRef.current = null;
}
setIsRunning(false);
}, [
accessToken,
selectedPromptIds,
selectedPolicies,
selectedGuardrails,
allFrameworks,
backendMode,
fixedModel,
requestProxyBaseUrl,
]);
const completedResults = testResults.filter((r) => r.status === "complete");
@@ -845,23 +928,30 @@ export default function ComplianceUI({
</div>
<div className="flex flex-col gap-1.5 pt-6 flex-shrink-0">
<button
type="button"
onClick={runTests}
disabled={selectedPromptIds.size === 0 || isRunning || disabledPersonalKeyCreation}
className={`flex items-center gap-1.5 px-4 py-2 rounded-lg text-sm font-medium transition-colors whitespace-nowrap ${selectedPromptIds.size === 0 || isRunning || disabledPersonalKeyCreation ? "bg-gray-100 text-gray-400 cursor-not-allowed" : "bg-blue-600 text-white hover:bg-blue-700"}`}
>
{isRunning ? (
<>
<Loader2 className="w-3.5 h-3.5 animate-spin" /> Running...
</>
) : (
<>
<Play className="w-3.5 h-3.5" /> Simulate (
{selectedPromptIds.size})
</>
)}
</button>
{isRunning ? (
<button
type="button"
onClick={() => batchAbortControllerRef.current?.abort()}
className="flex items-center gap-1.5 px-4 py-2 rounded-lg text-sm font-medium transition-colors whitespace-nowrap bg-red-600 text-white hover:bg-red-700"
>
<Square className="w-3.5 h-3.5" /> Stop
</button>
) : (
<button
type="button"
onClick={runTests}
disabled={selectedPromptIds.size === 0 || disabledPersonalKeyCreation}
className={`flex items-center gap-1.5 px-4 py-2 rounded-lg text-sm font-medium transition-colors whitespace-nowrap ${selectedPromptIds.size === 0 || disabledPersonalKeyCreation ? "bg-gray-100 text-gray-400 cursor-not-allowed" : "bg-blue-600 text-white hover:bg-blue-700"}`}
>
<Play className="w-3.5 h-3.5" /> Simulate (
{selectedPromptIds.size})
</button>
)}
{isRunning && (
<span className="text-[11px] text-gray-500 flex items-center gap-1">
<Loader2 className="w-3 h-3 animate-spin" /> Running...
</span>
)}
<button
type="button"
onClick={() => {
@@ -1624,6 +1714,16 @@ export default function ComplianceUI({
: "False positive — incorrectly blocked"}
</span>
</div>
{result.returnedText != null && result.returnedText !== "" && (
<div className="mt-1.5">
<span className="text-gray-400 block mb-0.5">
LLM response:
</span>
<div className="text-gray-700 bg-gray-50 rounded px-2 py-1.5 border border-gray-100 max-h-32 overflow-y-auto whitespace-pre-wrap break-words">
{result.returnedText}
</div>
</div>
)}
</div>
)}
</div>
@@ -1,6 +1,6 @@
// fetch_agents.tsx
import { getProxyBaseUrl, getGlobalLitellmHeaderName } from "../../networking";
import { getProxyBaseUrl, getGlobalLitellmHeaderName, modelInfoCall } from "../../networking";
export interface Agent {
agent_id: string;
@@ -13,6 +13,28 @@ export interface Agent {
};
}
/** MCP tool entry in the same format as chat completions API (litellm_params.tools) */
export interface MCPToolEntry {
type: "mcp";
server_label?: string;
server_url: string;
require_approval?: string;
allowed_tools?: string[];
}
/** Agent model from /model/info where litellm_params.model starts with "litellm_agent/" */
export interface AgentModel {
model_name: string;
litellm_params: {
model: string;
litellm_system_prompt?: string;
/** Saved MCP tools array (same shape as chat completions API tools) */
tools?: MCPToolEntry[];
[key: string]: unknown;
};
model_info?: Record<string, unknown> | null;
}
/**
* Fetches available A2A agents from /v1/agents endpoint.
*/
@@ -53,3 +75,44 @@ export const fetchAvailableAgents = async (
throw error;
}
};
/**
* Fetches available litellm_agent models from /v2/model/info.
* Filters for models where litellm_params.model starts with "litellm_agent/".
*/
export const fetchAvailableAgentModels = async (
accessToken: string,
userID: string,
userRole: string,
customBaseUrl?: string,
): Promise<AgentModel[]> => {
try {
const size = 200;
const response = await modelInfoCall(accessToken, userID, userRole, 1, size);
const data = response?.data ?? [];
const list = Array.isArray(data) ? data : [];
const agentModels: AgentModel[] = list
.filter(
(m: { litellm_params?: { model?: string } }) =>
typeof m?.litellm_params?.model === "string" &&
m.litellm_params.model.startsWith("litellm_agent/"),
)
.map((m: any) => ({
model_name: m.model_name ?? m.model_group ?? "",
litellm_params: {
...m.litellm_params,
model: m.litellm_params.model,
litellm_system_prompt: m.litellm_params?.litellm_system_prompt,
tools: Array.isArray(m.litellm_params?.tools) ? m.litellm_params.tools : undefined,
},
model_info: m.model_info ?? null,
}));
agentModels.sort((a, b) => a.model_name.localeCompare(b.model_name));
return agentModels;
} catch (error) {
console.error("Error fetching agent models:", error);
throw error;
}
};