mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_migrate_teams_page
# Conflicts: # ui/litellm-dashboard/e2e_tests/fixtures/migratedPages.ts # ui/litellm-dashboard/src/app/(dashboard)/page.tsx # ui/litellm-dashboard/src/utils/migratedPages.test.ts # ui/litellm-dashboard/src/utils/migratedPages.ts
This commit is contained in:
commit
6698d2f406
44 changed files with 7971 additions and 1862 deletions
49
.github/workflows/osv-scan.yml
vendored
Normal file
49
.github/workflows/osv-scan.yml
vendored
Normal file
|
|
@ -0,0 +1,49 @@
|
|||
name: OSV Scan
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_branch
|
||||
- "litellm_**"
|
||||
paths:
|
||||
- uv.lock
|
||||
- ui/litellm-dashboard/package-lock.json
|
||||
- osv-scanner.toml
|
||||
- .github/workflows/osv-scan.yml
|
||||
schedule:
|
||||
- cron: "23 6 * * *"
|
||||
workflow_dispatch:
|
||||
|
||||
permissions: {}
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
osv-scan:
|
||||
name: osv-scan
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
permissions:
|
||||
contents: read
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Download osv-scanner v2.3.8
|
||||
run: |
|
||||
curl -fsSL --retry 3 -o "$RUNNER_TEMP/osv-scanner" \
|
||||
https://github.com/google/osv-scanner/releases/download/v2.3.8/osv-scanner_linux_amd64
|
||||
echo "bc98e15319ed0d515e3f9235287ba53cdc5535d576d24fd573978ecfe9ab92dc $RUNNER_TEMP/osv-scanner" | sha256sum -c -
|
||||
chmod +x "$RUNNER_TEMP/osv-scanner"
|
||||
|
||||
- name: Scan lockfiles
|
||||
run: |
|
||||
"$RUNNER_TEMP/osv-scanner" scan source \
|
||||
--config osv-scanner.toml \
|
||||
-L uv.lock \
|
||||
-L ui/litellm-dashboard/package-lock.json
|
||||
|
|
@ -17,7 +17,7 @@ from litellm.integrations.otel.model.payloads import (
|
|||
ServiceSpanData,
|
||||
)
|
||||
from litellm.integrations.otel.plumbing.providers import to_otel_span_kind
|
||||
from litellm.integrations.otel.model.semconv import Error
|
||||
from litellm.integrations.otel.model.semconv import Error, ExceptionEvent
|
||||
from litellm.integrations.otel.model.spans import (
|
||||
SPAN_REGISTRY,
|
||||
SpanRole,
|
||||
|
|
@ -179,9 +179,17 @@ class SpanEmitter:
|
|||
else None
|
||||
)
|
||||
if error and (error.error_type or error.message):
|
||||
span.set_attribute(Error.TYPE, error.error_type or "error")
|
||||
span.set_status(
|
||||
Status(StatusCode.ERROR, error.message or error.error_type or "error")
|
||||
error_type = error.error_type or "error"
|
||||
message = error.message or error.error_type or "error"
|
||||
span.set_attribute(Error.TYPE, error_type)
|
||||
span.set_status(Status(StatusCode.ERROR, message))
|
||||
# Carry the full message on the standard ``exception`` event so backends
|
||||
# map it as full text under ``exception.message``. Setting it as a bare
|
||||
# string attribute instead lets backends like Elasticsearch dynamic-map
|
||||
# it to a ``keyword`` capped at 1024 chars, truncating the message.
|
||||
span.add_event(
|
||||
ExceptionEvent.NAME,
|
||||
{ExceptionEvent.TYPE: error_type, ExceptionEvent.MESSAGE: message},
|
||||
)
|
||||
# On success leave the status UNSET (the semconv default) rather than
|
||||
# forcing OK — that matches the FastAPI server span and avoids implying a
|
||||
|
|
|
|||
|
|
@ -146,6 +146,21 @@ class Error:
|
|||
TYPE: Final = "error.type"
|
||||
|
||||
|
||||
class ExceptionEvent:
|
||||
"""OTel exception-event name and attribute keys (semconv ``exception.*``).
|
||||
|
||||
The full error message rides ``exception.message`` on a span event rather than
|
||||
a custom string attribute. Backends recognise these semantic-convention names
|
||||
and map them as full text; an unrecognised key (e.g. ``error_message``) falls
|
||||
into the default dynamic template, which truncates strings to a 1024-char
|
||||
``keyword``.
|
||||
"""
|
||||
|
||||
NAME: Final = "exception"
|
||||
TYPE: Final = "exception.type"
|
||||
MESSAGE: Final = "exception.message"
|
||||
|
||||
|
||||
class Server:
|
||||
ADDRESS: Final = "server.address"
|
||||
PORT: Final = "server.port"
|
||||
|
|
|
|||
BIN
litellm/proxy/_experimental/out/assets/logos/cisco.png
Normal file
BIN
litellm/proxy/_experimental/out/assets/logos/cisco.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 1.9 KiB |
|
|
@ -0,0 +1,108 @@
|
|||
"""Cisco AI Defense Guardrail Integration for LiteLLM."""
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from litellm.types.guardrails import SupportedGuardrailIntegrations
|
||||
|
||||
from .cisco_ai_defense import (
|
||||
CiscoAIDefenseGuardrail,
|
||||
CiscoAIDefenseGuardrailAPIError,
|
||||
CiscoAIDefenseGuardrailMissingSecrets,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.guardrails import Guardrail, LitellmParams
|
||||
|
||||
|
||||
def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"):
|
||||
import litellm
|
||||
|
||||
guardrail_name = guardrail.get("guardrail_name")
|
||||
if not guardrail_name:
|
||||
raise ValueError("Cisco AI Defense: guardrail_name is required")
|
||||
|
||||
optional_params = getattr(litellm_params, "optional_params", None)
|
||||
|
||||
_callback = CiscoAIDefenseGuardrail(
|
||||
guardrail_name=guardrail_name,
|
||||
api_key=litellm_params.api_key,
|
||||
api_base=litellm_params.api_base,
|
||||
inspection_type=_get_optional_value(
|
||||
litellm_params, optional_params, "inspection_type"
|
||||
),
|
||||
inspect_path=_get_optional_value(
|
||||
litellm_params, optional_params, "inspect_path"
|
||||
),
|
||||
enabled_rules=_get_optional_value(
|
||||
litellm_params, optional_params, "enabled_rules"
|
||||
),
|
||||
integration_profile_id=_get_optional_value(
|
||||
litellm_params, optional_params, "integration_profile_id"
|
||||
),
|
||||
integration_profile_version=_get_optional_value(
|
||||
litellm_params, optional_params, "integration_profile_version"
|
||||
),
|
||||
integration_tenant_id=_get_optional_value(
|
||||
litellm_params, optional_params, "integration_tenant_id"
|
||||
),
|
||||
integration_type=_get_optional_value(
|
||||
litellm_params, optional_params, "integration_type"
|
||||
),
|
||||
on_flagged_action=_get_optional_value(
|
||||
litellm_params, optional_params, "on_flagged_action"
|
||||
),
|
||||
fallback_on_error=_get_optional_value(
|
||||
litellm_params, optional_params, "fallback_on_error"
|
||||
),
|
||||
timeout=_get_optional_value(litellm_params, optional_params, "timeout"),
|
||||
event_hook=litellm_params.mode,
|
||||
default_on=litellm_params.default_on or False,
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_callback)
|
||||
|
||||
# MCP post-tool-call hooks are dispatched through success callbacks.
|
||||
litellm.logging_callback_manager.add_litellm_success_callback(_callback)
|
||||
|
||||
return _callback
|
||||
|
||||
|
||||
def _get_optional_value(litellm_params, optional_params, attribute_name):
|
||||
"""Resolve Cisco optional params without inheriting sibling defaults."""
|
||||
if optional_params is not None:
|
||||
if isinstance(optional_params, dict):
|
||||
if attribute_name in optional_params:
|
||||
return optional_params[attribute_name]
|
||||
else:
|
||||
nested_fields_set = getattr(optional_params, "model_fields_set", None)
|
||||
if nested_fields_set is None or attribute_name in nested_fields_set:
|
||||
value = getattr(optional_params, attribute_name, None)
|
||||
if value is not None:
|
||||
return value
|
||||
|
||||
if litellm_params is None:
|
||||
return None
|
||||
# Only accept flattened values the caller explicitly set.
|
||||
fields_set = getattr(litellm_params, "model_fields_set", None)
|
||||
if fields_set is None or attribute_name not in fields_set:
|
||||
return None
|
||||
return getattr(litellm_params, attribute_name, None)
|
||||
|
||||
|
||||
guardrail_initializer_registry = {
|
||||
SupportedGuardrailIntegrations.CISCO_AI_DEFENSE.value: initialize_guardrail,
|
||||
}
|
||||
|
||||
|
||||
guardrail_class_registry = {
|
||||
SupportedGuardrailIntegrations.CISCO_AI_DEFENSE.value: CiscoAIDefenseGuardrail,
|
||||
}
|
||||
|
||||
|
||||
__all__ = [
|
||||
"CiscoAIDefenseGuardrail",
|
||||
"CiscoAIDefenseGuardrailAPIError",
|
||||
"CiscoAIDefenseGuardrailMissingSecrets",
|
||||
"initialize_guardrail",
|
||||
"guardrail_initializer_registry",
|
||||
"guardrail_class_registry",
|
||||
]
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -0,0 +1,704 @@
|
|||
"""MCP-specific inspection logic for the Cisco AI Defense guardrail.
|
||||
|
||||
The public guardrail class imports this private mixin from
|
||||
``cisco_ai_defense.py``. Keeping MCP logic here avoids circular imports
|
||||
while preserving the existing public import path.
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_to_applied_guardrails_header,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.mcp import MCPPostCallResponseObject
|
||||
|
||||
from .cisco_ai_defense import _ScanContext
|
||||
|
||||
|
||||
def _serialize_mcp_content_item(item: object) -> Dict[str, Any]:
|
||||
"""Serialize an MCP content item to a JSON-friendly dict.
|
||||
|
||||
Handles raw dicts, MCP SDK Pydantic models, and simple ``.text`` objects.
|
||||
"""
|
||||
if isinstance(item, dict):
|
||||
return dict(item)
|
||||
model_dump = getattr(item, "model_dump", None)
|
||||
if callable(model_dump):
|
||||
try:
|
||||
return dict(model_dump(exclude_none=True))
|
||||
except TypeError:
|
||||
return dict(model_dump())
|
||||
text = getattr(item, "text", None)
|
||||
if isinstance(text, str):
|
||||
return {"type": getattr(item, "type", "text"), "text": text}
|
||||
return {"type": "text", "text": str(item)}
|
||||
|
||||
|
||||
class _CiscoAIDefenseMcpMixin:
|
||||
"""MCP-specific instance methods for ``CiscoAIDefenseGuardrail``.
|
||||
|
||||
Holds the MCP hooks, JSON-RPC payload builders, and redaction helpers.
|
||||
"""
|
||||
|
||||
if TYPE_CHECKING:
|
||||
api_base: str
|
||||
inspect_path: str
|
||||
inspection_type: str
|
||||
_PROVIDER_NAME: str
|
||||
guardrail_name: Optional[str]
|
||||
|
||||
def should_run_guardrail(
|
||||
self, data: dict, event_type: GuardrailEventHooks
|
||||
) -> bool: ...
|
||||
|
||||
async def _post_inspection(
|
||||
self, url: str, payload: Dict[str, Any], surface: str
|
||||
) -> Dict[str, Any]: ...
|
||||
|
||||
def _handle_api_error(
|
||||
self,
|
||||
error: Exception,
|
||||
*,
|
||||
request_data: Optional[dict] = ...,
|
||||
start_time: Optional[datetime] = ...,
|
||||
surface: str = ...,
|
||||
direction: str = ...,
|
||||
) -> Dict[str, Any]: ...
|
||||
|
||||
def _finalize_inspection(
|
||||
self,
|
||||
inspect_response: Dict[str, Any],
|
||||
request_data: dict,
|
||||
context: "_ScanContext",
|
||||
start_time: datetime,
|
||||
response_obj: object = ...,
|
||||
) -> Dict[str, Any]: ...
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# MCP post-tool hook (dispatcher contract)
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def async_post_mcp_tool_call_hook(
|
||||
self,
|
||||
kwargs: dict,
|
||||
response_obj: "MCPPostCallResponseObject",
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
) -> Optional["MCPPostCallResponseObject"]:
|
||||
"""Scan MCP tool output and return a replacement object on block."""
|
||||
del start_time, end_time
|
||||
|
||||
if self.inspection_type != "mcp":
|
||||
return None
|
||||
|
||||
request_data: Dict[str, Any] = {}
|
||||
for key in (
|
||||
"name",
|
||||
"litellm_call_id",
|
||||
"id",
|
||||
"user",
|
||||
"mcp_tool_name",
|
||||
"tool_name",
|
||||
"mcp_arguments",
|
||||
"arguments",
|
||||
"mcp_server_name",
|
||||
"server_name",
|
||||
"metadata",
|
||||
"litellm_metadata",
|
||||
"mcp_tool_call_metadata",
|
||||
"guardrails",
|
||||
):
|
||||
if key in kwargs and kwargs[key] is not None:
|
||||
request_data[key] = kwargs[key]
|
||||
self._hydrate_mcp_tool_context(request_data)
|
||||
|
||||
if not (
|
||||
self.should_run_guardrail(
|
||||
data=request_data,
|
||||
event_type=GuardrailEventHooks.during_mcp_call,
|
||||
)
|
||||
or self.should_run_guardrail(
|
||||
data=request_data,
|
||||
event_type=GuardrailEventHooks.pre_mcp_call,
|
||||
)
|
||||
):
|
||||
verbose_proxy_logger.debug(
|
||||
"Cisco AI Defense guardrail (%s): no MCP mode configured "
|
||||
"— skipping MCP response scan.",
|
||||
self.guardrail_name,
|
||||
)
|
||||
return None
|
||||
|
||||
mcp_tool_response = self._extract_mcp_tool_call_response(response_obj)
|
||||
if mcp_tool_response is None:
|
||||
verbose_proxy_logger.debug(
|
||||
"Cisco AI Defense guardrail: no MCP tool response payload "
|
||||
"to scan, skipping"
|
||||
)
|
||||
return None
|
||||
|
||||
original_response = kwargs.get("original_response")
|
||||
try:
|
||||
await self._inspect_mcp_response(
|
||||
request_data=request_data,
|
||||
response=mcp_tool_response,
|
||||
redact_response_obj=(
|
||||
original_response
|
||||
if original_response is not None
|
||||
else mcp_tool_response
|
||||
),
|
||||
)
|
||||
except HTTPException as exc:
|
||||
blocking_response = self._build_blocking_mcp_response(
|
||||
detail=exc.detail, original_response_obj=response_obj
|
||||
)
|
||||
self._replace_mcp_tool_response(response_obj, blocking_response)
|
||||
if original_response is not None:
|
||||
self._replace_mcp_tool_response(original_response, blocking_response)
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=request_data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
verbose_proxy_logger.warning(
|
||||
"Cisco AI Defense guardrail (%s): MCP response blocked — "
|
||||
"tool output replaced with synthesized violation message.",
|
||||
self.guardrail_name,
|
||||
)
|
||||
return blocking_response
|
||||
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=request_data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
return None
|
||||
|
||||
def _build_blocking_mcp_response(
|
||||
self,
|
||||
detail: object,
|
||||
original_response_obj: object,
|
||||
) -> "MCPPostCallResponseObject":
|
||||
"""Build a synthetic MCPPostCallResponseObject for blocked output."""
|
||||
import json as _json
|
||||
|
||||
from litellm.types.llms.base import HiddenParams
|
||||
from litellm.types.mcp import MCPPostCallResponseObject
|
||||
from mcp.types import TextContent
|
||||
|
||||
if isinstance(detail, dict):
|
||||
payload = detail
|
||||
else:
|
||||
payload = {
|
||||
"error": "Blocked by Cisco AI Defense Guardrail",
|
||||
"message": (
|
||||
str(detail) if detail else "Blocked by Cisco AI Defense Guardrail"
|
||||
),
|
||||
"provider": self._PROVIDER_NAME,
|
||||
"guardrail": self.guardrail_name,
|
||||
"surface": "mcp",
|
||||
"direction": "output",
|
||||
"action": "block",
|
||||
}
|
||||
|
||||
original_hidden = getattr(original_response_obj, "hidden_params", None)
|
||||
if isinstance(original_hidden, HiddenParams):
|
||||
hidden_params: Any = original_hidden
|
||||
else:
|
||||
response_cost = getattr(original_hidden, "response_cost", None)
|
||||
hidden_params = (
|
||||
HiddenParams(response_cost=response_cost)
|
||||
if response_cost is not None
|
||||
else HiddenParams()
|
||||
)
|
||||
|
||||
return MCPPostCallResponseObject(
|
||||
mcp_tool_call_response=[
|
||||
TextContent(type="text", text=_json.dumps(payload))
|
||||
],
|
||||
hidden_params=hidden_params,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _replace_mcp_tool_response(
|
||||
response_obj: object, replacement_obj: object
|
||||
) -> bool:
|
||||
replacement = getattr(replacement_obj, "mcp_tool_call_response", None)
|
||||
if replacement is None:
|
||||
return False
|
||||
|
||||
inner = getattr(response_obj, "mcp_tool_call_response", None)
|
||||
if inner is not None:
|
||||
if _CiscoAIDefenseMcpMixin._replace_mcp_tool_response(
|
||||
inner, replacement_obj
|
||||
):
|
||||
return True
|
||||
try:
|
||||
setattr(response_obj, "mcp_tool_call_response", replacement)
|
||||
return True
|
||||
except (AttributeError, TypeError, ValueError):
|
||||
return False
|
||||
|
||||
content = getattr(response_obj, "content", None)
|
||||
if isinstance(content, list):
|
||||
content[:] = replacement
|
||||
structured_replacement = (
|
||||
_CiscoAIDefenseMcpMixin._replacement_structured_content(replacement)
|
||||
)
|
||||
if hasattr(response_obj, "structuredContent"):
|
||||
try:
|
||||
setattr(response_obj, "structuredContent", structured_replacement)
|
||||
except (AttributeError, TypeError, ValueError):
|
||||
pass
|
||||
if hasattr(response_obj, "isError"):
|
||||
try:
|
||||
setattr(response_obj, "isError", True)
|
||||
except (AttributeError, TypeError, ValueError):
|
||||
pass
|
||||
return True
|
||||
|
||||
if isinstance(response_obj, list):
|
||||
response_obj[:] = replacement
|
||||
return True
|
||||
|
||||
if isinstance(response_obj, dict):
|
||||
result = response_obj.get("result")
|
||||
if isinstance(result, dict):
|
||||
result["content"] = replacement
|
||||
result["structuredContent"] = (
|
||||
_CiscoAIDefenseMcpMixin._replacement_structured_content(replacement)
|
||||
)
|
||||
result["isError"] = True
|
||||
return True
|
||||
response_obj["result"] = {
|
||||
"content": replacement,
|
||||
"structuredContent": _CiscoAIDefenseMcpMixin._replacement_structured_content(
|
||||
replacement
|
||||
),
|
||||
"isError": True,
|
||||
}
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _replacement_structured_content(
|
||||
replacement: object,
|
||||
) -> Optional[Dict[str, str]]:
|
||||
if not isinstance(replacement, list) or not replacement:
|
||||
return None
|
||||
first = replacement[0]
|
||||
text = (
|
||||
first.get("text")
|
||||
if isinstance(first, dict)
|
||||
else getattr(first, "text", None)
|
||||
)
|
||||
return {"result": text} if isinstance(text, str) else None
|
||||
|
||||
@staticmethod
|
||||
def _extract_mcp_tool_call_response(response_obj: object) -> object:
|
||||
"""Pull the raw tool-call response off a MCPPostCallResponseObject."""
|
||||
inner = getattr(response_obj, "mcp_tool_call_response", None)
|
||||
if inner is None and isinstance(response_obj, dict):
|
||||
inner = response_obj.get("mcp_tool_call_response")
|
||||
return inner if inner is not None else response_obj
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# MCP request / response inspection
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _inspect_mcp_request(
|
||||
self,
|
||||
data: dict,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> Dict[str, Any]:
|
||||
del user_api_key_dict # carried via logging metadata, not the wire payload
|
||||
url = f"{self.api_base}{self.inspect_path}"
|
||||
payload = self._build_mcp_request_payload(data=data)
|
||||
if payload is None:
|
||||
verbose_proxy_logger.debug(
|
||||
"Cisco AI Defense guardrail: could not build MCP request "
|
||||
"payload, skipping"
|
||||
)
|
||||
return {}
|
||||
start_time = datetime.now()
|
||||
try:
|
||||
inspect_response = await self._post_inspection(
|
||||
url=url, payload=payload, surface="mcp"
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
return self._handle_api_error(
|
||||
exc,
|
||||
request_data=data,
|
||||
start_time=start_time,
|
||||
surface="mcp",
|
||||
direction="input",
|
||||
)
|
||||
|
||||
from .cisco_ai_defense import _ScanContext
|
||||
|
||||
return self._finalize_inspection(
|
||||
inspect_response=inspect_response,
|
||||
request_data=data,
|
||||
context=_ScanContext(surface="mcp", direction="input"),
|
||||
start_time=start_time,
|
||||
)
|
||||
|
||||
async def _inspect_mcp_response(
|
||||
self,
|
||||
request_data: dict,
|
||||
response: object,
|
||||
user_api_key_dict: Optional[UserAPIKeyAuth] = None,
|
||||
redact_response_obj: object = None,
|
||||
) -> Dict[str, Any]:
|
||||
del user_api_key_dict # carried via logging metadata, not the wire payload
|
||||
url = f"{self.api_base}{self.inspect_path}"
|
||||
payload = self._build_mcp_response_payload(
|
||||
request_data=request_data,
|
||||
response=response,
|
||||
)
|
||||
if payload is None:
|
||||
verbose_proxy_logger.debug(
|
||||
"Cisco AI Defense guardrail: could not build MCP response "
|
||||
"payload, skipping"
|
||||
)
|
||||
return {}
|
||||
start_time = datetime.now()
|
||||
try:
|
||||
inspect_response = await self._post_inspection(
|
||||
url=url, payload=payload, surface="mcp"
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as exc:
|
||||
return self._handle_api_error(
|
||||
exc,
|
||||
request_data=request_data,
|
||||
start_time=start_time,
|
||||
surface="mcp",
|
||||
direction="output",
|
||||
)
|
||||
|
||||
from .cisco_ai_defense import _ScanContext
|
||||
|
||||
return self._finalize_inspection(
|
||||
inspect_response=inspect_response,
|
||||
request_data=request_data,
|
||||
context=_ScanContext(surface="mcp", direction="output"),
|
||||
start_time=start_time,
|
||||
response_obj=(
|
||||
response if redact_response_obj is None else redact_response_obj
|
||||
),
|
||||
)
|
||||
|
||||
def _build_mcp_request_payload(
|
||||
self,
|
||||
data: dict,
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Build the JSON-RPC ``tools/call`` envelope sent to ``/inspect/mcp``.
|
||||
|
||||
The Cisco AI Defense MCP inspect endpoint expects the JSON-RPC
|
||||
envelope itself as the request body — *not* wrapped under a
|
||||
``request`` key with sibling ``metadata`` / ``config`` keys. Policies
|
||||
are applied based on the API key linked to the request. Operator
|
||||
metadata (user, call id, src/dst app, etc.) is carried out-of-band
|
||||
via the standard logging payload so the wire contract stays
|
||||
identical to a hand-rolled ``curl`` against ``/inspect/mcp``.
|
||||
"""
|
||||
if data.get("jsonrpc") == "2.0":
|
||||
return {
|
||||
"jsonrpc": "2.0",
|
||||
"id": (data.get("id") or data.get("litellm_call_id") or "litellm-mcp"),
|
||||
"method": data.get("method") or "tools/call",
|
||||
"params": data.get("params") or {},
|
||||
}
|
||||
|
||||
tool_name = (
|
||||
data.get("mcp_tool_name") or data.get("tool_name") or data.get("name")
|
||||
)
|
||||
if not tool_name:
|
||||
return None
|
||||
|
||||
arguments = data.get("mcp_arguments")
|
||||
if arguments is None:
|
||||
arguments = data.get("arguments")
|
||||
|
||||
return {
|
||||
"jsonrpc": "2.0",
|
||||
"id": data.get("litellm_call_id") or "litellm-mcp",
|
||||
"method": "tools/call",
|
||||
"params": {
|
||||
"name": tool_name,
|
||||
"arguments": (arguments if isinstance(arguments, dict) else {}),
|
||||
},
|
||||
}
|
||||
|
||||
def _build_mcp_response_payload(
|
||||
self,
|
||||
request_data: dict,
|
||||
response: object,
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Build the MCP response-inspection body sent to ``/inspect/mcp``."""
|
||||
request_payload = self._build_mcp_request_payload(data=request_data)
|
||||
if request_payload is None:
|
||||
return None
|
||||
normalized = self._normalize_mcp_response(response)
|
||||
if normalized is None:
|
||||
return None
|
||||
|
||||
payload = dict(request_payload)
|
||||
response_id = normalized.get("id")
|
||||
if response_id not in (None, "litellm-mcp"):
|
||||
payload["id"] = response_id
|
||||
elif payload.get("id") in (None, "litellm-mcp"):
|
||||
request_id = request_data.get("litellm_call_id") or request_data.get("id")
|
||||
if request_id:
|
||||
payload["id"] = request_id
|
||||
|
||||
if "result" in normalized:
|
||||
payload["result"] = normalized["result"]
|
||||
if "error" in normalized:
|
||||
payload["error"] = normalized["error"]
|
||||
return payload
|
||||
|
||||
@staticmethod
|
||||
def _hydrate_mcp_tool_context(request_data: Dict[str, Any]) -> None:
|
||||
metadata = request_data.get("mcp_tool_call_metadata")
|
||||
if metadata is None:
|
||||
nested = request_data.get("metadata") or request_data.get(
|
||||
"litellm_metadata"
|
||||
)
|
||||
if isinstance(nested, dict):
|
||||
metadata = nested.get("mcp_tool_call_metadata")
|
||||
if not isinstance(metadata, dict):
|
||||
return
|
||||
|
||||
name = metadata.get("name")
|
||||
arguments = metadata.get("arguments")
|
||||
server_name = metadata.get("mcp_server_name")
|
||||
|
||||
if name:
|
||||
request_data.setdefault("mcp_tool_name", name)
|
||||
request_data.setdefault("tool_name", name)
|
||||
request_data.setdefault("name", name)
|
||||
if arguments is not None:
|
||||
request_data.setdefault("mcp_arguments", arguments)
|
||||
request_data.setdefault("arguments", arguments)
|
||||
if server_name:
|
||||
request_data.setdefault("mcp_server_name", server_name)
|
||||
request_data.setdefault("server_name", server_name)
|
||||
|
||||
@staticmethod
|
||||
def _normalize_mcp_response(response: object) -> Optional[Dict[str, Any]]:
|
||||
"""Normalize an MCP tool response into a JSON-RPC envelope.
|
||||
|
||||
Handles JSON-RPC dicts, raw content lists, MCP SDK models, and
|
||||
Pydantic-coerced ``[(field_name, value)]`` lists.
|
||||
"""
|
||||
if isinstance(response, dict):
|
||||
if response.get("jsonrpc") == "2.0":
|
||||
return dict(response)
|
||||
if isinstance(response.get("result"), dict):
|
||||
return {
|
||||
"jsonrpc": "2.0",
|
||||
"id": response.get("id") or "litellm-mcp",
|
||||
"result": response["result"],
|
||||
}
|
||||
content = response.get("content")
|
||||
if isinstance(content, list):
|
||||
return {
|
||||
"jsonrpc": "2.0",
|
||||
"id": response.get("id") or "litellm-mcp",
|
||||
"result": _CiscoAIDefenseMcpMixin._build_mcp_result(
|
||||
content=content, source=response
|
||||
),
|
||||
}
|
||||
if isinstance(response, list):
|
||||
if response and all(
|
||||
isinstance(item, tuple) and len(item) == 2 and isinstance(item[0], str)
|
||||
for item in response
|
||||
):
|
||||
response_fields = dict(response)
|
||||
inner_content = response_fields.get("content")
|
||||
if isinstance(inner_content, list):
|
||||
return {
|
||||
"jsonrpc": "2.0",
|
||||
"id": "litellm-mcp",
|
||||
"result": _CiscoAIDefenseMcpMixin._build_mcp_result(
|
||||
content=inner_content, source=response_fields
|
||||
),
|
||||
}
|
||||
else:
|
||||
return None
|
||||
return {
|
||||
"jsonrpc": "2.0",
|
||||
"id": "litellm-mcp",
|
||||
"result": _CiscoAIDefenseMcpMixin._build_mcp_result(content=response),
|
||||
}
|
||||
model_dump = getattr(response, "model_dump", None)
|
||||
if callable(model_dump):
|
||||
try:
|
||||
dumped = model_dump(exclude_none=True)
|
||||
except TypeError:
|
||||
dumped = model_dump()
|
||||
if isinstance(dumped, dict):
|
||||
return _CiscoAIDefenseMcpMixin._normalize_mcp_response(dumped)
|
||||
content = getattr(response, "content", None)
|
||||
if isinstance(content, list):
|
||||
return {
|
||||
"jsonrpc": "2.0",
|
||||
"id": "litellm-mcp",
|
||||
"result": _CiscoAIDefenseMcpMixin._build_mcp_result(
|
||||
content=content, source=response
|
||||
),
|
||||
}
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _build_mcp_result(
|
||||
content: List[Any],
|
||||
source: object = None,
|
||||
) -> Dict[str, Any]:
|
||||
result: Dict[str, Any] = {
|
||||
"content": [_serialize_mcp_content_item(item) for item in content]
|
||||
}
|
||||
for key in ("structuredContent", "isError"):
|
||||
value = (
|
||||
source.get(key)
|
||||
if isinstance(source, dict)
|
||||
else getattr(source, key, None)
|
||||
)
|
||||
if value is not None and (key != "isError" or isinstance(value, bool)):
|
||||
result[key] = value
|
||||
return result
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# MCP redact (in-place rewrite of tool output)
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def _set_mcp_tool_response_text(response_obj: object, text: str) -> bool:
|
||||
"""Replace text content in any supported MCP response shape."""
|
||||
if response_obj is None:
|
||||
return False
|
||||
|
||||
inner = getattr(response_obj, "mcp_tool_call_response", None)
|
||||
if inner is not None:
|
||||
return _CiscoAIDefenseMcpMixin._set_mcp_tool_response_text(inner, text)
|
||||
|
||||
content_list = _CiscoAIDefenseMcpMixin._coerce_to_content_list(response_obj)
|
||||
|
||||
replaced = False
|
||||
if isinstance(content_list, list):
|
||||
for item in content_list:
|
||||
if isinstance(item, dict) and item.get("type") == "text":
|
||||
item["text"] = text
|
||||
replaced = True
|
||||
elif hasattr(item, "type") and getattr(item, "type", None) == "text":
|
||||
try:
|
||||
setattr(item, "text", text)
|
||||
replaced = True
|
||||
except (AttributeError, TypeError, ValueError):
|
||||
continue
|
||||
|
||||
replacement = {"result": text}
|
||||
if (
|
||||
isinstance(response_obj, list)
|
||||
and response_obj
|
||||
and all(
|
||||
isinstance(item, tuple) and len(item) == 2 and isinstance(item[0], str)
|
||||
for item in response_obj
|
||||
)
|
||||
):
|
||||
for index, item in enumerate(response_obj):
|
||||
if item[0] == "structuredContent":
|
||||
response_obj[index] = (item[0], replacement)
|
||||
replaced = True
|
||||
elif hasattr(response_obj, "structuredContent"):
|
||||
try:
|
||||
setattr(response_obj, "structuredContent", replacement)
|
||||
replaced = True
|
||||
except (AttributeError, TypeError, ValueError):
|
||||
pass
|
||||
elif isinstance(response_obj, dict):
|
||||
result = response_obj.get("result")
|
||||
target: Dict[Any, Any] = (
|
||||
result if isinstance(result, dict) else response_obj
|
||||
)
|
||||
if "structuredContent" in target:
|
||||
target["structuredContent"] = replacement
|
||||
replaced = True
|
||||
|
||||
return replaced
|
||||
|
||||
@staticmethod
|
||||
def _coerce_to_content_list(response_obj: object) -> Optional[List[Any]]:
|
||||
"""Find the MCP content list inside supported response shapes."""
|
||||
if response_obj is None:
|
||||
return None
|
||||
inner = getattr(response_obj, "mcp_tool_call_response", None)
|
||||
if inner is not None:
|
||||
return _CiscoAIDefenseMcpMixin._coerce_to_content_list(inner)
|
||||
content = getattr(response_obj, "content", None)
|
||||
if isinstance(content, list):
|
||||
return content
|
||||
if isinstance(response_obj, list):
|
||||
if response_obj and all(
|
||||
isinstance(item, tuple) and len(item) == 2 and isinstance(item[0], str)
|
||||
for item in response_obj
|
||||
):
|
||||
inner_content = dict(response_obj).get("content")
|
||||
if isinstance(inner_content, list):
|
||||
return inner_content
|
||||
return None
|
||||
return response_obj
|
||||
return None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# MCP-specific verdict extraction
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def _extract_sanitized_mcp_arguments(
|
||||
inspect_response: Dict[str, Any],
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""Pull sanitized MCP tool-call arguments off the verdict.
|
||||
|
||||
Cisco can return them at the top level (``params.arguments``) or
|
||||
under ``sanitized_payload`` / ``modified_payload``.
|
||||
"""
|
||||
containers = [inspect_response]
|
||||
for container_key in ("result", "data"):
|
||||
container = inspect_response.get(container_key)
|
||||
if isinstance(container, dict):
|
||||
containers.append(container)
|
||||
|
||||
for container in containers:
|
||||
params = container.get("params")
|
||||
if isinstance(params, dict):
|
||||
args = params.get("arguments")
|
||||
if isinstance(args, dict) and args:
|
||||
return dict(args)
|
||||
for key in (
|
||||
"sanitized_payload",
|
||||
"sanitizedPayload",
|
||||
"modified_payload",
|
||||
"modifiedPayload",
|
||||
):
|
||||
payload = container.get(key)
|
||||
if isinstance(payload, dict):
|
||||
inner_params = payload.get("params")
|
||||
if isinstance(inner_params, dict):
|
||||
args = inner_params.get("arguments")
|
||||
if isinstance(args, dict) and args:
|
||||
return dict(args)
|
||||
direct = payload.get("arguments")
|
||||
if isinstance(direct, dict) and direct:
|
||||
return dict(direct)
|
||||
return None
|
||||
|
|
@ -47,6 +47,9 @@ from litellm.types.proxy.guardrails.guardrail_hooks.qohash import (
|
|||
from litellm.types.proxy.guardrails.guardrail_hooks.vigil_guard import (
|
||||
VigilGuardGuardrailConfigModel,
|
||||
)
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.cisco_ai_defense import (
|
||||
CiscoAIDefenseGuardrailConfigModel,
|
||||
)
|
||||
|
||||
"""
|
||||
Pydantic object defining how to set guardrails on litellm proxy
|
||||
|
|
@ -80,6 +83,7 @@ class SupportedGuardrailIntegrations(Enum):
|
|||
PILLAR = "pillar"
|
||||
GRAYSWAN = "grayswan"
|
||||
PANW_PRISMA_AIRS = "panw_prisma_airs"
|
||||
CISCO_AI_DEFENSE = "cisco_ai_defense"
|
||||
AZURE_PROMPT_SHIELD = "azure/prompt_shield"
|
||||
AZURE_TEXT_MODERATIONS = "azure/text_moderations"
|
||||
MODEL_ARMOR = "model_armor"
|
||||
|
|
@ -840,6 +844,7 @@ class Mode(BaseModel):
|
|||
|
||||
|
||||
class LitellmParams(
|
||||
CiscoAIDefenseGuardrailConfigModel,
|
||||
PresidioConfigModel,
|
||||
BedrockGuardrailConfigModel,
|
||||
LakeraV2GuardrailConfigModel,
|
||||
|
|
|
|||
|
|
@ -0,0 +1,148 @@
|
|||
"""
|
||||
Cisco AI Defense Guardrail Config Model
|
||||
"""
|
||||
|
||||
from typing import List, Literal, Optional
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
from .base import GuardrailConfigModel
|
||||
|
||||
CISCO_AI_DEFENSE_RULE_NAMES = Literal[
|
||||
"Code Detection",
|
||||
"Harassment",
|
||||
"Hate Speech",
|
||||
"PCI",
|
||||
"PHI",
|
||||
"PII",
|
||||
"Prompt Injection",
|
||||
"Profanity",
|
||||
"Sexual Content & Exploitation",
|
||||
"Social Division & Polarization",
|
||||
"Violence & Public Safety Threats",
|
||||
]
|
||||
|
||||
|
||||
# Inspection surfaces supported by Cisco AI Defense. The Cisco Inspection API
|
||||
# exposes two separate endpoints — one for LLM chat conversations and one for
|
||||
# MCP tool calls. The user picks exactly one surface to scan per guardrail
|
||||
# instance; configure two guardrails if you need to scan both.
|
||||
CISCO_AI_DEFENSE_INSPECTION_TYPE = Literal["chat", "mcp"]
|
||||
|
||||
|
||||
class CiscoAIDefenseRule(BaseModel):
|
||||
"""A single rule to enable for Cisco AI Defense inspection."""
|
||||
|
||||
rule_name: CISCO_AI_DEFENSE_RULE_NAMES = Field(
|
||||
description="The canonical Cisco AI Defense rule name to evaluate.",
|
||||
)
|
||||
entity_types: Optional[List[str]] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Optional list of entity types for the rule (e.g. 'Email Address', "
|
||||
"'Phone Number'). Applies to rules such as PII, PCI, and PHI."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class CiscoAIDefenseGuardrailConfigModelOptionalParams(BaseModel):
|
||||
"""Optional parameters for the Cisco AI Defense guardrail."""
|
||||
|
||||
model_config = ConfigDict(extra="allow")
|
||||
|
||||
inspection_type: CISCO_AI_DEFENSE_INSPECTION_TYPE = Field(
|
||||
default="chat",
|
||||
description=(
|
||||
"Which Cisco AI Defense inspection surface to use. "
|
||||
"'chat' scans LLM model conversations via /api/v1/inspect/chat. "
|
||||
"'mcp' scans MCP tool calls via /api/v1/inspect/mcp. "
|
||||
"Each guardrail instance targets exactly one surface; configure "
|
||||
"two guardrails to scan both chat and MCP traffic."
|
||||
),
|
||||
)
|
||||
inspect_path: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Override for the inspection endpoint path. Defaults to "
|
||||
"/api/v1/inspect/chat when inspection_type='chat' and "
|
||||
"/api/v1/inspect/mcp when inspection_type='mcp'."
|
||||
),
|
||||
)
|
||||
enabled_rules: Optional[List[CiscoAIDefenseRule]] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Explicit list of Cisco AI Defense rules to evaluate. If omitted, "
|
||||
"the policies configured for the API key in the Cisco AI Defense "
|
||||
"UI are used."
|
||||
),
|
||||
)
|
||||
integration_profile_id: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Integration profile id to apply (advanced).",
|
||||
)
|
||||
integration_profile_version: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Integration profile version to apply (advanced).",
|
||||
)
|
||||
integration_tenant_id: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Integration tenant id to apply (advanced).",
|
||||
)
|
||||
integration_type: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Integration type to apply (advanced).",
|
||||
)
|
||||
on_flagged_action: Optional[str] = Field(
|
||||
default="block",
|
||||
description=(
|
||||
"Action to take when Cisco AI Defense flags content. 'block' raises "
|
||||
"an HTTPException; 'monitor' logs the detection and lets the "
|
||||
"request continue."
|
||||
),
|
||||
)
|
||||
fallback_on_error: Optional[Literal["allow", "block"]] = Field(
|
||||
default="block",
|
||||
description=(
|
||||
"Behaviour when the Cisco AI Defense API is unavailable: 'allow' "
|
||||
"proceeds without scanning (high availability), 'block' rejects "
|
||||
"the request (maximum security)."
|
||||
),
|
||||
)
|
||||
timeout: Optional[float] = Field(
|
||||
default=10.0,
|
||||
ge=1.0,
|
||||
le=60.0,
|
||||
description="Timeout (seconds) for Cisco AI Defense API calls (1-60).",
|
||||
)
|
||||
|
||||
|
||||
class CiscoAIDefenseGuardrailConfigModel(
|
||||
GuardrailConfigModel[CiscoAIDefenseGuardrailConfigModelOptionalParams]
|
||||
):
|
||||
"""Configuration parameters for the Cisco AI Defense guardrail."""
|
||||
|
||||
api_key: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"API key for the Cisco AI Defense inspection endpoint. If "
|
||||
"not provided, the `CISCO_AI_DEFENSE_API_KEY` environment variable "
|
||||
"is used. Sent in the `X-Cisco-AI-Defense-API-Key` header. "
|
||||
"Both the chat and MCP endpoints use this key."
|
||||
),
|
||||
)
|
||||
api_base: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Regional base URL for the Cisco AI Defense Inspection API. "
|
||||
"Defaults to https://us.api.inspect.aidefense.security.cisco.com. "
|
||||
"Supported regions: us (us-west-2), ap (ap-ne-1), eu "
|
||||
"(eu-central-1). The environment variable "
|
||||
"`CISCO_AI_DEFENSE_API_BASE` is consulted as a fallback. The "
|
||||
"endpoint path is derived from inspection_type "
|
||||
"(/api/v1/inspect/chat for 'chat', /api/v1/inspect/mcp for 'mcp')."
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def ui_friendly_name() -> str:
|
||||
return "Cisco AI Defense"
|
||||
14
osv-scanner.toml
Normal file
14
osv-scanner.toml
Normal file
|
|
@ -0,0 +1,14 @@
|
|||
[[IgnoredVulns]]
|
||||
id = "GHSA-w8v5-vhqr-4h9v"
|
||||
ignoreUntil = 2026-09-09
|
||||
reason = "diskcache has no fixed release published; remove this entry once one exists"
|
||||
|
||||
[[IgnoredVulns]]
|
||||
id = "GHSA-hg6j-4rv6-33pg"
|
||||
ignoreUntil = 2026-08-15
|
||||
reason = "aiohttp held at 3.13.5: vcrpy releases <= 8.1.1 cannot import aiohttp >= 3.14 and the merged upstream fix (vcrpy PR 996) is unreleased; bump aiohttp and drop this entry when a newer vcrpy ships"
|
||||
|
||||
[[IgnoredVulns]]
|
||||
id = "GHSA-jg22-mg44-37j8"
|
||||
ignoreUntil = 2026-08-15
|
||||
reason = "aiohttp held at 3.13.5: vcrpy releases <= 8.1.1 cannot import aiohttp >= 3.14 and the merged upstream fix (vcrpy PR 996) is unreleased; bump aiohttp and drop this entry when a newer vcrpy ships"
|
||||
|
|
@ -141,6 +141,12 @@ ANTHROPIC_DIRECT_MODELS: Tuple[ModelEntry, ...] = (
|
|||
mode="adaptive",
|
||||
required_env=_ANTHROPIC_REQ,
|
||||
caps=_CAPS_XHIGH_MAX,
|
||||
fail_reason=(
|
||||
"claude-fable-5 is not yet released on the Anthropic API for the CI "
|
||||
"account; Anthropic returns not_found_error until the model is "
|
||||
"available, so this cell stays loud in CI. Remove this fail_reason "
|
||||
"once the model is available."
|
||||
),
|
||||
),
|
||||
ModelEntry(
|
||||
alias="claude-opus-4-8",
|
||||
|
|
|
|||
|
|
@ -70,6 +70,11 @@ def test_map_response_format():
|
|||
assert result == {"response_format": response_format}
|
||||
|
||||
|
||||
_AUDIO_FILE_PATH = os.path.join(
|
||||
os.path.dirname(os.path.realpath(__file__)), "gettysburg.wav"
|
||||
)
|
||||
|
||||
|
||||
class TestFireworksAIAudioTranscription(BaseLLMAudioTranscriptionTest):
|
||||
def get_base_audio_transcription_call_args(self) -> dict:
|
||||
return {
|
||||
|
|
@ -80,6 +85,60 @@ class TestFireworksAIAudioTranscription(BaseLLMAudioTranscriptionTest):
|
|||
def get_custom_llm_provider(self) -> litellm.LlmProviders:
|
||||
return litellm.LlmProviders.FIREWORKS_AI
|
||||
|
||||
def test_audio_transcription(self):
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from openai.types.audio import Transcription
|
||||
|
||||
audio_file = open(_AUDIO_FILE_PATH, "rb")
|
||||
mock_client = MagicMock()
|
||||
mock_client.audio.transcriptions.create.return_value = Transcription(
|
||||
text="four score and seven years ago"
|
||||
)
|
||||
|
||||
transcript = transcription(
|
||||
**self.get_base_audio_transcription_call_args(),
|
||||
file=audio_file,
|
||||
api_key="fw-test-key",
|
||||
client=mock_client,
|
||||
)
|
||||
|
||||
assert transcript.text == "four score and seven years ago"
|
||||
sent = mock_client.audio.transcriptions.create.call_args.kwargs
|
||||
assert sent["model"] == "whisper-v3"
|
||||
assert sent["file"] is audio_file
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_audio_transcription_async(self):
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from openai.types.audio import Transcription
|
||||
|
||||
audio_file = open(_AUDIO_FILE_PATH, "rb")
|
||||
raw_response = MagicMock()
|
||||
raw_response.headers = {}
|
||||
raw_response.parse.return_value = Transcription(
|
||||
text="four score and seven years ago"
|
||||
)
|
||||
mock_client = MagicMock()
|
||||
mock_client.audio.transcriptions.with_raw_response.create = AsyncMock(
|
||||
return_value=raw_response
|
||||
)
|
||||
|
||||
transcript = await litellm.atranscription(
|
||||
**self.get_base_audio_transcription_call_args(),
|
||||
file=audio_file,
|
||||
api_key="fw-test-key",
|
||||
client=mock_client,
|
||||
)
|
||||
|
||||
assert transcript.text == "four score and seven years ago"
|
||||
sent = (
|
||||
mock_client.audio.transcriptions.with_raw_response.create.call_args.kwargs
|
||||
)
|
||||
assert sent["model"] == "whisper-v3"
|
||||
assert sent["file"] is audio_file
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"disable_add_transform_inline_image_block",
|
||||
|
|
|
|||
|
|
@ -502,6 +502,90 @@ def test_emitter_without_call_id_is_not_deduped():
|
|||
assert len(exporter.get_finished_spans()) == 2
|
||||
|
||||
|
||||
def _emit_error_span(message, error_type="litellm.APIError"):
|
||||
from litellm.integrations.otel.emitter import SpanEmitter
|
||||
|
||||
cfg = OpenTelemetryV2Config(exporter="in_memory")
|
||||
provider, exporter = providers.in_memory_provider(cfg)
|
||||
engine = SpanEmitter(providers.get_tracer(provider, "t"), cfg)
|
||||
data = LLMCallSpanData(
|
||||
operation=GenAIOperation.CHAT,
|
||||
provider="openai",
|
||||
request_model="gpt-4o",
|
||||
response_model=None,
|
||||
response_id=None,
|
||||
request_params=LLMRequestParams(),
|
||||
usage=LLMUsage(),
|
||||
finish_reasons=(),
|
||||
error=SpanError(error_type=error_type, message=message),
|
||||
response_cost=None,
|
||||
server=None,
|
||||
identity=RequestIdentity(call_id=None),
|
||||
)
|
||||
engine.emit(SpanRole.LLM_CALL, data)
|
||||
(span,) = exporter.get_finished_spans()
|
||||
return span
|
||||
|
||||
|
||||
def _exception_event(span):
|
||||
from litellm.integrations.otel.model.semconv import ExceptionEvent
|
||||
|
||||
events = [e for e in span.events if e.name == ExceptionEvent.NAME]
|
||||
assert len(events) == 1, "expected exactly one exception event"
|
||||
return events[0]
|
||||
|
||||
|
||||
def test_error_message_recorded_as_full_exception_event_untruncated():
|
||||
"""Regression for the Elasticsearch keyword/ignore_above:1024 truncation.
|
||||
|
||||
A long error message must survive intact on the standard ``exception``
|
||||
event under ``exception.message`` — not get dropped onto a bare string
|
||||
attribute that backends dynamic-map to a 1024-char ``keyword``. The SDK
|
||||
must not truncate it either, so a 5000-char message stays 5000 chars.
|
||||
"""
|
||||
from litellm.integrations.otel.model.semconv import Error, ExceptionEvent
|
||||
|
||||
long_message = "boom: " + "x" * 5000
|
||||
span = _emit_error_span(long_message, error_type="litellm.APIError")
|
||||
|
||||
event = _exception_event(span)
|
||||
assert event.attributes[ExceptionEvent.MESSAGE] == long_message
|
||||
assert len(event.attributes[ExceptionEvent.MESSAGE]) == len(long_message) > 1024
|
||||
assert event.attributes[ExceptionEvent.TYPE] == "litellm.APIError"
|
||||
|
||||
# error.type stays a low-cardinality attribute; the message does NOT become a
|
||||
# bare string attribute (which is what got truncated).
|
||||
assert span.attributes[Error.TYPE] == "litellm.APIError"
|
||||
assert ExceptionEvent.MESSAGE not in span.attributes
|
||||
assert span.status.description == long_message
|
||||
|
||||
|
||||
def test_success_span_records_no_exception_event():
|
||||
from litellm.integrations.otel.emitter import SpanEmitter
|
||||
from litellm.integrations.otel.model.semconv import ExceptionEvent
|
||||
|
||||
cfg = OpenTelemetryV2Config(exporter="in_memory")
|
||||
provider, exporter = providers.in_memory_provider(cfg)
|
||||
engine = SpanEmitter(providers.get_tracer(provider, "t"), cfg)
|
||||
data = LLMCallSpanData(
|
||||
operation=GenAIOperation.CHAT,
|
||||
provider="openai",
|
||||
request_model="gpt-4o",
|
||||
response_model="gpt-4o",
|
||||
response_id="resp-1",
|
||||
request_params=LLMRequestParams(),
|
||||
usage=LLMUsage(),
|
||||
finish_reasons=("stop",),
|
||||
error=None,
|
||||
response_cost=None,
|
||||
server=None,
|
||||
identity=RequestIdentity(call_id=None),
|
||||
)
|
||||
engine.emit(SpanRole.LLM_CALL, data)
|
||||
(span,) = exporter.get_finished_spans()
|
||||
assert all(e.name != ExceptionEvent.NAME for e in span.events)
|
||||
|
||||
|
||||
# --- service taxonomy: which calls become spans, and of what kind ----------- #
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,362 @@
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
from contextlib import contextmanager
|
||||
from datetime import datetime
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, Dict
|
||||
from unittest.mock import AsyncMock, patch
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
from httpx import Request, Response
|
||||
from litellm.types.utils import (
|
||||
Choices,
|
||||
Delta,
|
||||
Message,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
StreamingChoices,
|
||||
TextChoices,
|
||||
TextCompletionResponse,
|
||||
)
|
||||
|
||||
|
||||
def _make_text_completion_response(text: str) -> TextCompletionResponse:
|
||||
return TextCompletionResponse(
|
||||
choices=[{"text": text, "index": 0, "finish_reason": "stop"}]
|
||||
)
|
||||
|
||||
|
||||
def _make_model_response_with_content(content: str) -> ModelResponse:
|
||||
return ModelResponse(
|
||||
choices=[
|
||||
Choices(
|
||||
index=0,
|
||||
finish_reason="stop",
|
||||
message=Message(role="assistant", content=content),
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
import litellm
|
||||
from litellm import DualCache
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails.guardrail_hooks.cisco_ai_defense import (
|
||||
CiscoAIDefenseGuardrail,
|
||||
CiscoAIDefenseGuardrailMissingSecrets,
|
||||
)
|
||||
from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
|
||||
|
||||
CISCO_BASE = "https://us.api.inspect.aidefense.security.cisco.com"
|
||||
CHAT_URL = f"{CISCO_BASE}/api/v1/inspect/chat"
|
||||
MCP_URL = f"{CISCO_BASE}/api/v1/inspect/mcp"
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _patch_inspection_post(g: CiscoAIDefenseGuardrail, post_mock: Any):
|
||||
async def _send(request: Request, **kwargs: Any) -> Response:
|
||||
return await post_mock(
|
||||
url=str(request.url),
|
||||
headers=request.headers,
|
||||
json=json.loads(request.content.decode("utf-8")),
|
||||
follow_redirects=kwargs.get("follow_redirects"),
|
||||
)
|
||||
|
||||
with patch.object(g.async_handler.client, "send", new=_send):
|
||||
yield post_mock
|
||||
|
||||
|
||||
def _mock_inspect_response(
|
||||
json_body: dict, *, status: int = 200, url: str = CHAT_URL
|
||||
) -> Response:
|
||||
return Response(
|
||||
status_code=status,
|
||||
json=json_body,
|
||||
request=Request(method="POST", url=url),
|
||||
)
|
||||
|
||||
|
||||
def _safe_response(url: str = CHAT_URL) -> Response:
|
||||
return _mock_inspect_response(
|
||||
{
|
||||
"is_safe": True,
|
||||
"classifications": [],
|
||||
"severity": "NONE_SEVERITY",
|
||||
"rules": [],
|
||||
"action": "allow",
|
||||
},
|
||||
url=url,
|
||||
)
|
||||
|
||||
|
||||
def _violation_response(url: str = CHAT_URL) -> Response:
|
||||
return _mock_inspect_response(
|
||||
{
|
||||
"is_safe": False,
|
||||
"classifications": ["SECURITY_VIOLATION", "PRIVACY_VIOLATION"],
|
||||
"severity": "HIGH",
|
||||
"rules": [
|
||||
{"rule_name": "Prompt Injection"},
|
||||
{"rule_name": "PII", "entity_types": ["Email Address"]},
|
||||
],
|
||||
"explanation": "Detected jailbreak attempt with PII exfiltration",
|
||||
"event_id": "evt_123",
|
||||
"action": "block",
|
||||
},
|
||||
url=url,
|
||||
)
|
||||
|
||||
|
||||
def _mcp_request(name="lookup", args=None, jsonrpc=False, **extra):
|
||||
args = args if args is not None else {}
|
||||
if jsonrpc:
|
||||
return {
|
||||
"jsonrpc": "2.0",
|
||||
"id": "1",
|
||||
"method": "tools/call",
|
||||
"params": {"name": name, "arguments": args},
|
||||
**extra,
|
||||
}
|
||||
return {"mcp_tool_name": name, "mcp_arguments": args, **extra}
|
||||
|
||||
|
||||
def _mcp_response(content=None, response_cost=0.0):
|
||||
if content is None:
|
||||
content = [{"type": "text", "text": "ok"}]
|
||||
return SimpleNamespace(
|
||||
mcp_tool_call_response=content,
|
||||
hidden_params=SimpleNamespace(response_cost=response_cost),
|
||||
)
|
||||
|
||||
|
||||
def _mcp_result_text(content) -> str:
|
||||
if not content:
|
||||
return ""
|
||||
item = content[0] if isinstance(content, list) else content
|
||||
return getattr(item, "text", None) or item.get("text", "")
|
||||
|
||||
|
||||
def _chat_request_tool_call_args(arguments: str) -> dict:
|
||||
return {
|
||||
"messages": [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_1",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "send_data",
|
||||
"arguments": arguments,
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
def _chat_request_function_call_args(arguments: str) -> dict:
|
||||
return {
|
||||
"messages": [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"function_call": {
|
||||
"name": "exfil",
|
||||
"arguments": arguments,
|
||||
},
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
def _redact_response(
|
||||
*,
|
||||
sanitized_text=None,
|
||||
sanitized_messages=None,
|
||||
sanitized_mcp_arguments=None,
|
||||
sanitized_payload=None,
|
||||
classifications=("PRIVACY_VIOLATION",),
|
||||
rules=({"rule_name": "PII"},),
|
||||
severity="HIGH",
|
||||
url=CHAT_URL,
|
||||
):
|
||||
body = {
|
||||
"is_safe": False,
|
||||
"classifications": list(classifications),
|
||||
"severity": severity,
|
||||
"rules": list(rules),
|
||||
"action": "redact",
|
||||
}
|
||||
if sanitized_text is not None:
|
||||
body["sanitized_text"] = sanitized_text
|
||||
if sanitized_messages is not None:
|
||||
body["sanitized_messages"] = sanitized_messages
|
||||
if sanitized_mcp_arguments is not None:
|
||||
body["sanitized_mcp_arguments"] = sanitized_mcp_arguments
|
||||
if sanitized_payload is not None:
|
||||
body["sanitized_payload"] = sanitized_payload
|
||||
return _mock_inspect_response(body, url=url)
|
||||
|
||||
|
||||
def _responses_api_response(text, role="assistant"):
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.responses.main import GenericResponseOutputItem, OutputText
|
||||
|
||||
return ResponsesAPIResponse(
|
||||
id="resp_1",
|
||||
created_at=0,
|
||||
output=[
|
||||
GenericResponseOutputItem(
|
||||
type="message",
|
||||
id="msg_1",
|
||||
status="completed",
|
||||
role=role,
|
||||
content=[OutputText(type="output_text", text=text, annotations=[])],
|
||||
)
|
||||
],
|
||||
parallel_tool_calls=False,
|
||||
tool_choice=None,
|
||||
tools=None,
|
||||
top_p=None,
|
||||
usage=None,
|
||||
)
|
||||
|
||||
|
||||
def _make_guardrail(
|
||||
inspection_type="chat",
|
||||
event_hook="pre_call",
|
||||
*,
|
||||
name="t",
|
||||
api_key="x",
|
||||
default_on=True,
|
||||
**kwargs,
|
||||
):
|
||||
return CiscoAIDefenseGuardrail(
|
||||
guardrail_name=name,
|
||||
api_key=api_key,
|
||||
inspection_type=inspection_type,
|
||||
event_hook=event_hook,
|
||||
default_on=default_on,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
def _find_callback(name):
|
||||
from litellm.proxy.guardrails.guardrail_hooks.cisco_ai_defense import (
|
||||
CiscoAIDefenseGuardrail,
|
||||
)
|
||||
|
||||
for cb in litellm.callbacks:
|
||||
if isinstance(cb, CiscoAIDefenseGuardrail) and cb.guardrail_name == name:
|
||||
return cb
|
||||
raise AssertionError(f"Cisco guardrail {name!r} not in litellm.callbacks")
|
||||
|
||||
|
||||
def _make_streaming_chunks(parts):
|
||||
chunks = []
|
||||
for i, part in enumerate(parts):
|
||||
chunks.append(
|
||||
ModelResponseStream(
|
||||
id="resp_1",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
delta=Delta(content=part, role="assistant" if i == 0 else None),
|
||||
finish_reason="stop" if i == len(parts) - 1 else None,
|
||||
index=0,
|
||||
)
|
||||
],
|
||||
created=1234567890,
|
||||
model="gpt-4",
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
)
|
||||
return chunks
|
||||
|
||||
|
||||
async def _aiter(items):
|
||||
for item in items:
|
||||
yield item
|
||||
|
||||
|
||||
async def _streaming_setup(
|
||||
g,
|
||||
chunks,
|
||||
cisco_response=None,
|
||||
upstream=None,
|
||||
request_data=None,
|
||||
post_mock=None,
|
||||
):
|
||||
if post_mock is None:
|
||||
post_mock = (
|
||||
AsyncMock(return_value=cisco_response) if cisco_response else AsyncMock()
|
||||
)
|
||||
stream_source = upstream if upstream is not None else _aiter(chunks)
|
||||
if request_data is None:
|
||||
request_data = {"messages": [{"role": "user", "content": "hi"}]}
|
||||
received: list = []
|
||||
with _patch_inspection_post(g, post_mock):
|
||||
async for chunk in g.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
response=stream_source,
|
||||
request_data=request_data,
|
||||
):
|
||||
received.append(chunk)
|
||||
return received, post_mock
|
||||
|
||||
|
||||
__all__ = [
|
||||
"Any",
|
||||
"AsyncMock",
|
||||
"CHAT_URL",
|
||||
"CISCO_BASE",
|
||||
"Choices",
|
||||
"CiscoAIDefenseGuardrail",
|
||||
"CiscoAIDefenseGuardrailMissingSecrets",
|
||||
"Delta",
|
||||
"Dict",
|
||||
"DualCache",
|
||||
"HTTPException",
|
||||
"MCP_URL",
|
||||
"Message",
|
||||
"ModelResponse",
|
||||
"ModelResponseStream",
|
||||
"Request",
|
||||
"Response",
|
||||
"SimpleNamespace",
|
||||
"StreamingChoices",
|
||||
"TextChoices",
|
||||
"TextCompletionResponse",
|
||||
"UserAPIKeyAuth",
|
||||
"_aiter",
|
||||
"_chat_request_function_call_args",
|
||||
"_chat_request_tool_call_args",
|
||||
"_find_callback",
|
||||
"_make_guardrail",
|
||||
"_make_model_response_with_content",
|
||||
"_make_streaming_chunks",
|
||||
"_make_text_completion_response",
|
||||
"_mcp_request",
|
||||
"_mcp_response",
|
||||
"_mcp_result_text",
|
||||
"_mock_inspect_response",
|
||||
"_patch_inspection_post",
|
||||
"_redact_response",
|
||||
"_responses_api_response",
|
||||
"_safe_response",
|
||||
"_streaming_setup",
|
||||
"_violation_response",
|
||||
"contextmanager",
|
||||
"datetime",
|
||||
"init_guardrails_v2",
|
||||
"json",
|
||||
"litellm",
|
||||
"os",
|
||||
"patch",
|
||||
"pytest",
|
||||
"sys",
|
||||
]
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -0,0 +1,832 @@
|
|||
from tests.test_litellm.proxy.guardrails.guardrail_hooks._cisco_ai_defense_test_utils import (
|
||||
Any,
|
||||
AsyncMock,
|
||||
CiscoAIDefenseGuardrail,
|
||||
Dict,
|
||||
DualCache,
|
||||
HTTPException,
|
||||
MCP_URL,
|
||||
Response,
|
||||
SimpleNamespace,
|
||||
UserAPIKeyAuth,
|
||||
_make_guardrail,
|
||||
_make_model_response_with_content,
|
||||
_mcp_request,
|
||||
_mcp_response,
|
||||
_mcp_result_text,
|
||||
_mock_inspect_response,
|
||||
_patch_inspection_post,
|
||||
_redact_response,
|
||||
_safe_response,
|
||||
_violation_response,
|
||||
datetime,
|
||||
init_guardrails_v2,
|
||||
json,
|
||||
litellm,
|
||||
pytest,
|
||||
)
|
||||
|
||||
|
||||
def test_cisco_ai_defense_config_via_init_v2_mcp(monkeypatch):
|
||||
monkeypatch.setenv("CISCO_AI_DEFENSE_API_KEY", "test-key")
|
||||
litellm.guardrail_name_config_map = {}
|
||||
|
||||
init_guardrails_v2(
|
||||
all_guardrails=[
|
||||
{
|
||||
"guardrail_name": "cisco-mcp",
|
||||
"litellm_params": {
|
||||
"guardrail": "cisco_ai_defense",
|
||||
"mode": "pre_mcp_call",
|
||||
"default_on": True,
|
||||
"optional_params": {"inspection_type": "mcp"},
|
||||
},
|
||||
}
|
||||
],
|
||||
config_file_path="",
|
||||
)
|
||||
|
||||
|
||||
class TestCiscoAIDefenseMCPMode:
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_mode_inspects_mcp_request(self):
|
||||
g = _make_guardrail(inspection_type="mcp", event_hook="pre_mcp_call")
|
||||
data = _mcp_request(
|
||||
name="send_email", args={"to": "x@y.com"}, litellm_call_id="call-1"
|
||||
)
|
||||
post_mock = AsyncMock(return_value=_safe_response(url=MCP_URL))
|
||||
with _patch_inspection_post(g, post_mock):
|
||||
result = await g.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type="mcp_call",
|
||||
)
|
||||
assert result == data
|
||||
assert post_mock.call_args.kwargs["url"] == MCP_URL
|
||||
assert post_mock.call_args.kwargs["follow_redirects"] is False
|
||||
sent_payload = post_mock.call_args.kwargs["json"]
|
||||
assert sent_payload["jsonrpc"] == "2.0"
|
||||
assert sent_payload["method"] == "tools/call"
|
||||
assert sent_payload["params"]["name"] == "send_email"
|
||||
assert sent_payload["params"]["arguments"] == {"to": "x@y.com"}
|
||||
assert "request" not in sent_payload
|
||||
assert "metadata" not in sent_payload
|
||||
assert "config" not in sent_payload
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_mode_blocks_violation(self):
|
||||
g = _make_guardrail(inspection_type="mcp", event_hook="pre_mcp_call")
|
||||
data = _mcp_request(name="leak_secrets", args={"target": "evil"})
|
||||
with _patch_inspection_post(
|
||||
g, AsyncMock(return_value=_violation_response(url=MCP_URL))
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await g.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type="mcp_call",
|
||||
)
|
||||
assert exc.value.detail["surface"] == "mcp"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_mode_skips_chat_traffic(self):
|
||||
g = _make_guardrail(inspection_type="mcp", event_hook="pre_mcp_call")
|
||||
data = {"messages": [{"role": "user", "content": "hello"}]}
|
||||
post_mock = AsyncMock()
|
||||
with _patch_inspection_post(g, post_mock):
|
||||
result = await g.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type="completion",
|
||||
)
|
||||
assert result == data
|
||||
post_mock.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_mode_inspects_jsonrpc_envelope(self):
|
||||
g = _make_guardrail(inspection_type="mcp", event_hook="pre_mcp_call")
|
||||
data = _mcp_request(name="do_thing", args={"x": 1}, jsonrpc=True, id="abc")
|
||||
post_mock = AsyncMock(return_value=_safe_response(url=MCP_URL))
|
||||
with _patch_inspection_post(g, post_mock):
|
||||
await g.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type="mcp_call",
|
||||
)
|
||||
sent_payload = post_mock.call_args.kwargs["json"]
|
||||
assert sent_payload["jsonrpc"] == "2.0"
|
||||
assert sent_payload["id"] == "abc"
|
||||
assert sent_payload["params"]["name"] == "do_thing"
|
||||
assert sent_payload["params"]["arguments"] == {"x": 1}
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"verdict_extra",
|
||||
[
|
||||
{"sanitized_payload": {"params": {"arguments": {"note": "ssn [REDACTED]"}}}},
|
||||
{"sanitized_text": "ssn [REDACTED]"},
|
||||
],
|
||||
ids=["structured_arguments", "sanitized_text_fallback"],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_input_redaction_reaches_tool_call(self, verdict_extra):
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.utils import ProxyLogging
|
||||
|
||||
original_args = {"note": "ssn 123-45-6789"}
|
||||
sanitized_args = {"note": "ssn [REDACTED]"}
|
||||
|
||||
g = _make_guardrail(
|
||||
inspection_type="mcp",
|
||||
event_hook="pre_mcp_call",
|
||||
on_flagged_action="monitor",
|
||||
)
|
||||
data = _mcp_request(name="send_email", args=dict(original_args))
|
||||
cisco_resp = _mock_inspect_response(
|
||||
{
|
||||
"is_safe": False,
|
||||
"classifications": ["PRIVACY_VIOLATION"],
|
||||
"severity": "HIGH",
|
||||
"rules": [{"rule_name": "PII"}],
|
||||
"action": "redact",
|
||||
**verdict_extra,
|
||||
},
|
||||
url=MCP_URL,
|
||||
)
|
||||
|
||||
with _patch_inspection_post(g, AsyncMock(return_value=cisco_resp)):
|
||||
result = await g.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type="mcp_call",
|
||||
)
|
||||
|
||||
forwarded = ProxyLogging(
|
||||
user_api_key_cache=UserApiKeyCache()
|
||||
)._convert_mcp_hook_response_to_kwargs(
|
||||
response_data=result, original_kwargs={"arguments": dict(original_args)}
|
||||
)
|
||||
assert forwarded["arguments"] == sanitized_args, (
|
||||
"Sanitized MCP arguments did not reach the tool call. The proxy "
|
||||
"bridge forwards redactions only via ``modified_arguments``, so a "
|
||||
"redact verdict proceeded while the original unsanitized arguments "
|
||||
f"still hit the MCP server. Got: {forwarded['arguments']!r}"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_response_hook_inspects_tool_output(self):
|
||||
g = _make_guardrail(
|
||||
inspection_type="mcp", event_hook=["pre_mcp_call", "during_mcp_call"]
|
||||
)
|
||||
|
||||
response_obj = _mcp_response(
|
||||
SimpleNamespace(
|
||||
content=[{"type": "text", "text": "Here is the secret API key abc123"}]
|
||||
)
|
||||
)
|
||||
|
||||
post_mock = AsyncMock(return_value=_safe_response(url=MCP_URL))
|
||||
kwargs = {
|
||||
"name": "lookup_secret",
|
||||
"arguments": {"key": "production"},
|
||||
"mcp_server_name": "vault",
|
||||
"litellm_call_id": "call-42",
|
||||
}
|
||||
with _patch_inspection_post(g, post_mock):
|
||||
result = await g.async_post_mcp_tool_call_hook(
|
||||
kwargs=kwargs,
|
||||
response_obj=response_obj,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
|
||||
assert result is None
|
||||
assert post_mock.called
|
||||
assert post_mock.call_args.kwargs["url"] == MCP_URL
|
||||
sent_payload = post_mock.call_args.kwargs["json"]
|
||||
assert sent_payload["jsonrpc"] == "2.0"
|
||||
assert sent_payload["id"] == "call-42"
|
||||
assert sent_payload["method"] == "tools/call"
|
||||
assert sent_payload["params"] == {
|
||||
"name": "lookup_secret",
|
||||
"arguments": {"key": "production"},
|
||||
}
|
||||
assert sent_payload["result"]["content"][0]["text"] == (
|
||||
"Here is the secret API key abc123"
|
||||
)
|
||||
assert "request" not in sent_payload
|
||||
assert "metadata" not in sent_payload
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_response_hook_blocks_violation(self):
|
||||
from litellm.types.mcp import MCPPostCallResponseObject
|
||||
|
||||
g = _make_guardrail(
|
||||
inspection_type="mcp", event_hook=["pre_mcp_call", "during_mcp_call"]
|
||||
)
|
||||
response_obj = _mcp_response(
|
||||
SimpleNamespace(content=[{"type": "text", "text": "leaked"}])
|
||||
)
|
||||
|
||||
post_mock = AsyncMock(return_value=_violation_response(url=MCP_URL))
|
||||
with _patch_inspection_post(g, post_mock):
|
||||
result = await g.async_post_mcp_tool_call_hook(
|
||||
kwargs={"name": "leak", "arguments": {}},
|
||||
response_obj=response_obj,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
|
||||
assert result is not None, (
|
||||
"MCP response block was silently dropped — the litellm "
|
||||
"dispatcher swallows raised exceptions, so the hook must "
|
||||
"return a non-None MCPPostCallResponseObject to enforce a block."
|
||||
)
|
||||
assert isinstance(result, MCPPostCallResponseObject)
|
||||
replacement = result.mcp_tool_call_response
|
||||
assert len(replacement) == 1
|
||||
text = _mcp_result_text(replacement)
|
||||
assert "Blocked by Cisco AI Defense" in text
|
||||
assert "evt_123" in text
|
||||
assert "SECURITY_VIOLATION" in text
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_response_hook_skipped_in_chat_mode(self):
|
||||
g = _make_guardrail()
|
||||
response_obj = _mcp_response(
|
||||
SimpleNamespace(content=[{"type": "text", "text": "hi"}])
|
||||
)
|
||||
|
||||
post_mock = AsyncMock()
|
||||
with _patch_inspection_post(g, post_mock):
|
||||
result = await g.async_post_mcp_tool_call_hook(
|
||||
kwargs={"name": "tool", "arguments": {}},
|
||||
response_obj=response_obj,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
assert result is None
|
||||
post_mock.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_skipped_for_mcp_mode_guardrail(self):
|
||||
g = _make_guardrail(inspection_type="mcp", event_hook="pre_mcp_call")
|
||||
data = {"messages": [{"role": "user", "content": "hi"}]}
|
||||
response = _make_model_response_with_content("fine")
|
||||
|
||||
post_mock = AsyncMock()
|
||||
with _patch_inspection_post(g, post_mock):
|
||||
result = await g.async_post_call_success_hook(
|
||||
data=data,
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
response=response,
|
||||
)
|
||||
assert result is response
|
||||
post_mock.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_response_hook_runs_with_pre_mcp_call_only(self):
|
||||
g = _make_guardrail(inspection_type="mcp", event_hook="pre_mcp_call")
|
||||
response_obj = _mcp_response(
|
||||
SimpleNamespace(
|
||||
content=[{"type": "text", "text": "would have been scanned"}]
|
||||
)
|
||||
)
|
||||
|
||||
post_mock = AsyncMock(return_value=_safe_response(url=MCP_URL))
|
||||
with _patch_inspection_post(g, post_mock):
|
||||
await g.async_post_mcp_tool_call_hook(
|
||||
kwargs={"name": "lookup", "arguments": {}},
|
||||
response_obj=response_obj,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
|
||||
assert post_mock.called, (
|
||||
"MCP response scan was skipped when only ``pre_mcp_call`` "
|
||||
"was configured. Per product decision, pre_mcp_call means "
|
||||
"'guard the MCP call' — request AND response."
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"cisco_response_kind,expected_block",
|
||||
[("safe", False), ("violation", True)],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_response_hook_handles_raw_list_content(
|
||||
self, cisco_response_kind, expected_block
|
||||
):
|
||||
from litellm.types.mcp import MCPPostCallResponseObject
|
||||
|
||||
g = _make_guardrail(
|
||||
inspection_type="mcp", event_hook=["pre_mcp_call", "during_mcp_call"]
|
||||
)
|
||||
|
||||
text_content = (
|
||||
"exfiltrated data: ..."
|
||||
if cisco_response_kind == "violation"
|
||||
else "Here is the secret API key abc123"
|
||||
)
|
||||
response_obj = _mcp_response([{"type": "text", "text": text_content}])
|
||||
|
||||
cisco_resp = (
|
||||
_violation_response(url=MCP_URL)
|
||||
if cisco_response_kind == "violation"
|
||||
else _safe_response(url=MCP_URL)
|
||||
)
|
||||
post_mock = AsyncMock(return_value=cisco_resp)
|
||||
kwargs = {
|
||||
"name": "leak" if expected_block else "lookup_secret",
|
||||
"arguments": {"key": "production"} if not expected_block else {},
|
||||
"mcp_server_name": "vault",
|
||||
"litellm_call_id": "call-raw-list",
|
||||
}
|
||||
with _patch_inspection_post(g, post_mock):
|
||||
result = await g.async_post_mcp_tool_call_hook(
|
||||
kwargs=kwargs,
|
||||
response_obj=response_obj,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
|
||||
assert post_mock.called, (
|
||||
"MCP response inspect was silently skipped for raw-list "
|
||||
"shape — _normalize_mcp_response failed."
|
||||
)
|
||||
assert post_mock.call_args.kwargs["url"] == MCP_URL
|
||||
|
||||
if expected_block:
|
||||
assert isinstance(result, MCPPostCallResponseObject)
|
||||
replacement = result.mcp_tool_call_response
|
||||
assert len(replacement) == 1
|
||||
assert "Blocked by Cisco AI Defense" in _mcp_result_text(replacement)
|
||||
else:
|
||||
sent_payload = post_mock.call_args.kwargs["json"]
|
||||
assert sent_payload["jsonrpc"] == "2.0"
|
||||
assert sent_payload["id"] == "call-raw-list"
|
||||
assert sent_payload["method"] == "tools/call"
|
||||
assert sent_payload["params"] == {
|
||||
"name": "lookup_secret",
|
||||
"arguments": {"key": "production"},
|
||||
}
|
||||
assert sent_payload["result"]["content"][0]["text"] == text_content
|
||||
assert result is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_response_hook_through_real_logging_wrapper(self):
|
||||
from mcp.types import CallToolResult, TextContent
|
||||
|
||||
from litellm.types.mcp import MCPPostCallResponseObject
|
||||
|
||||
g = _make_guardrail(
|
||||
inspection_type="mcp", event_hook=["pre_mcp_call", "during_mcp_call"]
|
||||
)
|
||||
|
||||
real_result = CallToolResult(
|
||||
content=[TextContent(type="text", text="leak 9045629876")],
|
||||
structuredContent={"patient": {"ssn": "123-45-6789"}},
|
||||
isError=False,
|
||||
)
|
||||
wrapped = MCPPostCallResponseObject(
|
||||
mcp_tool_call_response=real_result,
|
||||
hidden_params={},
|
||||
)
|
||||
|
||||
assert isinstance(wrapped.mcp_tool_call_response, list)
|
||||
assert all(
|
||||
isinstance(item, tuple) and len(item) == 2
|
||||
for item in wrapped.mcp_tool_call_response
|
||||
), (
|
||||
"Pydantic coercion shape changed — update the normalizer to "
|
||||
"match the new wire format."
|
||||
)
|
||||
|
||||
post_mock = AsyncMock(return_value=_safe_response(url=MCP_URL))
|
||||
with _patch_inspection_post(g, post_mock):
|
||||
result = await g.async_post_mcp_tool_call_hook(
|
||||
kwargs={
|
||||
"name": "leak_tool",
|
||||
"arguments": {},
|
||||
"mcp_server_name": "vault",
|
||||
"litellm_call_id": "real-wire-call",
|
||||
},
|
||||
response_obj=wrapped,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
|
||||
assert post_mock.called, (
|
||||
"Inspect API not called for real CallToolResult shape — "
|
||||
"_normalize_mcp_response failed to handle Pydantic's "
|
||||
"iterated-BaseModel coercion."
|
||||
)
|
||||
assert post_mock.call_args.kwargs["url"] == MCP_URL
|
||||
sent_payload = post_mock.call_args.kwargs["json"]
|
||||
content_items = sent_payload["result"]["content"]
|
||||
|
||||
assert len(content_items) == 1, (
|
||||
f"expected exactly 1 content item from the real "
|
||||
f"CallToolResult.content list, got {len(content_items)}: "
|
||||
f"{content_items!r}"
|
||||
)
|
||||
assert content_items[0].get("text") == "leak 9045629876", (
|
||||
f"Cisco wire payload missed the real tool text; got "
|
||||
f"{content_items[0]!r}. This means the Pydantic-coerced "
|
||||
f"(field_name, value) tuple shape was serialized as text "
|
||||
f"content instead of being unwrapped to find the inner "
|
||||
f"``content`` field."
|
||||
)
|
||||
assert content_items[0].get("type") == "text"
|
||||
assert sent_payload["result"]["structuredContent"] == {
|
||||
"patient": {"ssn": "123-45-6789"}
|
||||
}
|
||||
assert sent_payload["result"]["isError"] is False
|
||||
assert sent_payload["id"] == "real-wire-call"
|
||||
assert sent_payload["method"] == "tools/call"
|
||||
assert sent_payload["params"] == {"name": "leak_tool", "arguments": {}}
|
||||
assert result is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_response_hook_uses_standard_logging_tool_metadata(self):
|
||||
g = _make_guardrail(inspection_type="mcp", event_hook="pre_mcp_call")
|
||||
response_obj = _mcp_response([{"type": "text", "text": "tool output"}])
|
||||
|
||||
post_mock = AsyncMock(return_value=_safe_response(url=MCP_URL))
|
||||
with _patch_inspection_post(g, post_mock):
|
||||
result = await g.async_post_mcp_tool_call_hook(
|
||||
kwargs={
|
||||
"litellm_call_id": "metadata-call",
|
||||
"mcp_tool_call_metadata": {
|
||||
"name": "lookup_secret",
|
||||
"arguments": {"key": "production"},
|
||||
"mcp_server_name": "vault",
|
||||
},
|
||||
},
|
||||
response_obj=response_obj,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
|
||||
assert result is None
|
||||
sent_payload = post_mock.call_args.kwargs["json"]
|
||||
assert sent_payload["method"] == "tools/call"
|
||||
assert sent_payload["params"] == {
|
||||
"name": "lookup_secret",
|
||||
"arguments": {"key": "production"},
|
||||
}
|
||||
assert sent_payload["result"]["content"][0]["text"] == "tool output"
|
||||
|
||||
|
||||
class TestCiscoAIDefenseRedactListShape:
|
||||
|
||||
@staticmethod
|
||||
def _violation_with_redact_response(text: str = "[REDACTED tool output]"):
|
||||
return _mock_inspect_response(
|
||||
{
|
||||
"is_safe": False,
|
||||
"classifications": ["PRIVACY_VIOLATION"],
|
||||
"severity": "HIGH",
|
||||
"rules": [{"rule_name": "PII", "entity_types": ["SSN"]}],
|
||||
"explanation": "PII detected, redaction available",
|
||||
"event_id": "evt_redact_1",
|
||||
"action": "redact",
|
||||
"sanitized_text": text,
|
||||
},
|
||||
url=MCP_URL,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _raw_list_factory():
|
||||
original_content = [{"type": "text", "text": "Your SSN is 123-45-6789."}]
|
||||
return original_content, lambda: original_content[0]["text"]
|
||||
|
||||
@staticmethod
|
||||
def _pydantic_tuple_list_factory():
|
||||
from mcp.types import TextContent
|
||||
|
||||
inner_content = [TextContent(type="text", text="SSN: 123-45-6789")]
|
||||
tuples_list = [
|
||||
("meta", None),
|
||||
("content", inner_content),
|
||||
("structuredContent", {"patient": {"ssn": "123-45-6789"}}),
|
||||
("isError", False),
|
||||
]
|
||||
return tuples_list, lambda: inner_content[0].text
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"factory_name",
|
||||
["_raw_list_factory", "_pydantic_tuple_list_factory"],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_redact_rewrites_mcp_response_list_shape(self, factory_name):
|
||||
|
||||
from litellm.types.mcp import MCPPostCallResponseObject
|
||||
|
||||
g = _make_guardrail(
|
||||
inspection_type="mcp", event_hook=["pre_mcp_call", "during_mcp_call"]
|
||||
)
|
||||
|
||||
content, get_text = getattr(self, factory_name)()
|
||||
response_obj = _mcp_response(content)
|
||||
|
||||
with _patch_inspection_post(
|
||||
g, AsyncMock(return_value=self._violation_with_redact_response())
|
||||
):
|
||||
result = await g.async_post_mcp_tool_call_hook(
|
||||
kwargs={"name": "leak", "arguments": {}},
|
||||
response_obj=response_obj,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
|
||||
assert result is None or not isinstance(result, MCPPostCallResponseObject), (
|
||||
f"Redact silently fell through to block for {factory_name}. "
|
||||
f"result={result!r}"
|
||||
)
|
||||
assert get_text() == "[REDACTED tool output]", (
|
||||
f"Redact silently failed for {factory_name}; original text "
|
||||
f"not rewritten."
|
||||
)
|
||||
if factory_name == "_pydantic_tuple_list_factory":
|
||||
structured_content = dict(content)["structuredContent"]
|
||||
assert structured_content == {"result": "[REDACTED tool output]"}
|
||||
assert "123-45-6789" not in json.dumps(structured_content)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_redact_rewrites_client_visible_original_response(self):
|
||||
from mcp.types import CallToolResult, TextContent
|
||||
|
||||
from litellm.types.llms.base import HiddenParams
|
||||
from litellm.types.mcp import MCPPostCallResponseObject
|
||||
|
||||
original_response = CallToolResult(
|
||||
content=[TextContent(type="text", text="SSN: 123-45-6789")],
|
||||
structuredContent={"patient": {"ssn": "123-45-6789"}},
|
||||
isError=False,
|
||||
)
|
||||
wrapper = MCPPostCallResponseObject(
|
||||
mcp_tool_call_response=original_response,
|
||||
hidden_params=HiddenParams(),
|
||||
)
|
||||
|
||||
g = _make_guardrail(
|
||||
inspection_type="mcp", event_hook=["pre_mcp_call", "during_mcp_call"]
|
||||
)
|
||||
with _patch_inspection_post(
|
||||
g, AsyncMock(return_value=self._violation_with_redact_response())
|
||||
):
|
||||
await g.async_post_mcp_tool_call_hook(
|
||||
kwargs={
|
||||
"name": "leak",
|
||||
"arguments": {},
|
||||
"original_response": original_response,
|
||||
},
|
||||
response_obj=wrapper,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
|
||||
assert original_response.content[0].text == "[REDACTED tool output]"
|
||||
assert "123-45-6789" not in json.dumps(original_response.structuredContent), (
|
||||
"Redact verdict left the client-visible MCP tool output unchanged. "
|
||||
"The post-call hook receives a wrapped MCPPostCallResponseObject but "
|
||||
"the endpoint returns kwargs['original_response'], so the redaction "
|
||||
"must rewrite that object too. structuredContent still leaks: "
|
||||
f"{original_response.structuredContent!r}"
|
||||
)
|
||||
|
||||
|
||||
class TestCiscoAIDefenseMcpInputRedactionFallback:
|
||||
"""``sanitized_text``-only redaction of structured MCP arguments."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_single_string_arg_is_rewritten(self):
|
||||
g = _make_guardrail(inspection_type="mcp", event_hook="pre_mcp_call")
|
||||
data = _mcp_request(
|
||||
name="search", args={"query": "my SSN is 123-45-6789", "limit": 10}
|
||||
)
|
||||
cisco = _redact_response(sanitized_text="my SSN is [REDACTED]", url=MCP_URL)
|
||||
with _patch_inspection_post(g, AsyncMock(return_value=cisco)):
|
||||
result = await g.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type="mcp_call",
|
||||
)
|
||||
assert result == data
|
||||
assert data["mcp_arguments"]["query"] == "my SSN is [REDACTED]"
|
||||
assert data["mcp_arguments"]["limit"] == 10
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ambiguous_multi_string_args_block_instead_of_leaking(self):
|
||||
g = _make_guardrail(
|
||||
inspection_type="mcp",
|
||||
event_hook="pre_mcp_call",
|
||||
on_flagged_action="block",
|
||||
)
|
||||
original = {"query": "PII data", "filter": "sensitive term", "limit": 10}
|
||||
data = _mcp_request(name="search", args=dict(original))
|
||||
cisco = _redact_response(sanitized_text="[REDACTED]", url=MCP_URL)
|
||||
with _patch_inspection_post(g, AsyncMock(return_value=cisco)):
|
||||
with pytest.raises(HTTPException):
|
||||
await g.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type="mcp_call",
|
||||
)
|
||||
assert data["mcp_arguments"] == original
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ambiguous_multi_string_args_not_partially_redacted_in_monitor(self):
|
||||
g = _make_guardrail(
|
||||
inspection_type="mcp",
|
||||
event_hook="pre_mcp_call",
|
||||
on_flagged_action="monitor",
|
||||
)
|
||||
original = {"query": "PII data", "filter": "sensitive term"}
|
||||
data = _mcp_request(name="search", args=dict(original))
|
||||
cisco = _redact_response(sanitized_text="[REDACTED]", url=MCP_URL)
|
||||
with _patch_inspection_post(g, AsyncMock(return_value=cisco)):
|
||||
result = await g.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type="mcp_call",
|
||||
)
|
||||
assert result == data
|
||||
assert data["mcp_arguments"] == original
|
||||
|
||||
|
||||
class TestCiscoAIDefenseMCPBlockingContract:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_block_response_survives_dispatcher_contract(self):
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.types.mcp import MCPPostCallResponseObject
|
||||
from mcp.types import CallToolResult, TextContent
|
||||
|
||||
g = _make_guardrail(
|
||||
name="cisco-mcp",
|
||||
inspection_type="mcp",
|
||||
event_hook=["pre_mcp_call", "during_mcp_call"],
|
||||
)
|
||||
raw_response = CallToolResult(
|
||||
content=[TextContent(type="text", text="exfiltrated")],
|
||||
structuredContent={"result": "exfiltrated"},
|
||||
isError=False,
|
||||
)
|
||||
response_obj = MCPPostCallResponseObject(
|
||||
mcp_tool_call_response=raw_response,
|
||||
hidden_params={},
|
||||
)
|
||||
|
||||
post_mock = AsyncMock(return_value=_violation_response(url=MCP_URL))
|
||||
captured: Dict[str, Any] = {}
|
||||
with _patch_inspection_post(g, post_mock):
|
||||
try:
|
||||
captured["result"] = await g.async_post_mcp_tool_call_hook(
|
||||
kwargs={
|
||||
"name": "leak",
|
||||
"arguments": {},
|
||||
"original_response": raw_response,
|
||||
},
|
||||
response_obj=response_obj,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
except Exception as e:
|
||||
captured["swallowed"] = repr(e)
|
||||
|
||||
assert "swallowed" not in captured, (
|
||||
f"async_post_mcp_tool_call_hook raised — the litellm "
|
||||
f"dispatcher would swallow this and the block would be lost. "
|
||||
f"Got: {captured.get('swallowed')}"
|
||||
)
|
||||
result = captured["result"]
|
||||
assert isinstance(result, MCPPostCallResponseObject), (
|
||||
"Hook must keep returning a MCPPostCallResponseObject for "
|
||||
"dispatcher paths that do honor returned replacements."
|
||||
)
|
||||
assert raw_response.isError is True
|
||||
assert "Blocked by Cisco AI Defense" in raw_response.content[0].text
|
||||
assert raw_response.structuredContent is not None
|
||||
assert "Blocked by Cisco AI Defense" in raw_response.structuredContent["result"]
|
||||
assert "exfiltrated" not in raw_response.structuredContent["result"]
|
||||
logging_stub = Logging.__new__(Logging)
|
||||
logging_stub.model_call_details = {}
|
||||
parsed = logging_stub._parse_post_mcp_call_hook_response(response=result)
|
||||
assert parsed is not None
|
||||
assert "Blocked by Cisco AI Defense" in _mcp_result_text(parsed)
|
||||
|
||||
|
||||
class TestCiscoAIDefenseJsonRpcSuccessEnvelope:
|
||||
|
||||
@staticmethod
|
||||
def _cisco_mcp_envelope(*, is_safe: bool, action: str = "Block") -> Response:
|
||||
return _mock_inspect_response(
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": 3,
|
||||
"result": {
|
||||
"is_safe": is_safe,
|
||||
"action": action,
|
||||
"classifications": [],
|
||||
"rules": [
|
||||
{
|
||||
"rule_name": "PII",
|
||||
"rule_id": 0,
|
||||
"entity_types": [],
|
||||
"classification": "NONE_VIOLATION",
|
||||
}
|
||||
],
|
||||
"event_id": "645d9d22-b016-47e0-a12c-9d587fb11c57",
|
||||
"detected_pii": [],
|
||||
},
|
||||
},
|
||||
url=MCP_URL,
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"is_safe,action,should_block",
|
||||
[
|
||||
(False, "Block", True),
|
||||
(True, "Allow", False),
|
||||
(False, "Allow", False),
|
||||
(True, "Block", True),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_jsonrpc_envelope_respects_verdict(
|
||||
self, is_safe, action, should_block
|
||||
):
|
||||
g = _make_guardrail(
|
||||
name="cisco-mcp", inspection_type="mcp", event_hook="pre_mcp_call"
|
||||
)
|
||||
data = _mcp_request(
|
||||
name="ask_question",
|
||||
args={
|
||||
"repoName": "facebook/react",
|
||||
"question": "What is React Fiber 9045629876?",
|
||||
},
|
||||
)
|
||||
with _patch_inspection_post(
|
||||
g,
|
||||
AsyncMock(
|
||||
return_value=self._cisco_mcp_envelope(is_safe=is_safe, action=action)
|
||||
),
|
||||
):
|
||||
if should_block:
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
await g.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type="mcp_call",
|
||||
)
|
||||
assert exc.value.status_code == 400
|
||||
assert exc.value.detail["surface"] == "mcp"
|
||||
assert (
|
||||
exc.value.detail["event_id"]
|
||||
== "645d9d22-b016-47e0-a12c-9d587fb11c57"
|
||||
)
|
||||
else:
|
||||
result = await g.async_pre_call_hook(
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
cache=DualCache(),
|
||||
data=data,
|
||||
call_type="mcp_call",
|
||||
)
|
||||
assert result == data
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"verdict,expected",
|
||||
[
|
||||
(
|
||||
{
|
||||
"is_safe": False,
|
||||
"classifications": ["SECURITY_VIOLATION"],
|
||||
"action": "block",
|
||||
},
|
||||
"passthrough",
|
||||
),
|
||||
(
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": 1,
|
||||
"result": {"is_safe": False, "action": "Block"},
|
||||
},
|
||||
{"is_safe": False, "action": "Block"},
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_unwrap_verdict_envelope(self, verdict, expected):
|
||||
unwrapped = CiscoAIDefenseGuardrail._unwrap_verdict_envelope(verdict)
|
||||
if expected == "passthrough":
|
||||
assert unwrapped is verdict
|
||||
else:
|
||||
assert unwrapped == expected
|
||||
|
|
@ -41,6 +41,7 @@ export const MIGRATED_E2E_PAGES: Record<string, string> = {
|
|||
"router-settings": "router-settings",
|
||||
users: "users",
|
||||
teams: "teams",
|
||||
organizations: "organizations",
|
||||
};
|
||||
|
||||
export const MIGRATED_E2E_SEGMENTS: string[] = [...new Set(Object.values(MIGRATED_E2E_PAGES))];
|
||||
|
|
|
|||
|
|
@ -678,14 +678,6 @@
|
|||
"count": 1
|
||||
}
|
||||
},
|
||||
"src/components/WebRTCTester.jsx": {
|
||||
"no-restricted-syntax": {
|
||||
"count": 2
|
||||
},
|
||||
"react/no-unescaped-entities": {
|
||||
"count": 2
|
||||
}
|
||||
},
|
||||
"src/components/activity_metrics.tsx": {
|
||||
"no-restricted-imports": {
|
||||
"count": 1
|
||||
|
|
|
|||
214
ui/litellm-dashboard/package-lock.json
generated
214
ui/litellm-dashboard/package-lock.json
generated
|
|
@ -751,9 +751,9 @@
|
|||
"license": "MIT"
|
||||
},
|
||||
"node_modules/@esbuild/aix-ppc64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/aix-ppc64/-/aix-ppc64-0.27.7.tgz",
|
||||
"integrity": "sha512-EKX3Qwmhz1eMdEJokhALr0YiD0lhQNwDqkPYyPhiSwKrh7/4KRjQc04sZ8db+5DVVnZ1LmbNDI1uAMPEUBnQPg==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/aix-ppc64/-/aix-ppc64-0.28.1.tgz",
|
||||
"integrity": "sha512-Svl7tq8k/08+p6CXPpRjQ1fKX+1odH/BQbb48fV6fj3CWHhsoIOoY87w1oHXm0qEpkIK3ZfVgp0hed3XBXzXMQ==",
|
||||
"cpu": [
|
||||
"ppc64"
|
||||
],
|
||||
|
|
@ -768,9 +768,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/android-arm": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/android-arm/-/android-arm-0.27.7.tgz",
|
||||
"integrity": "sha512-jbPXvB4Yj2yBV7HUfE2KHe4GJX51QplCN1pGbYjvsyCZbQmies29EoJbkEc+vYuU5o45AfQn37vZlyXy4YJ8RQ==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/android-arm/-/android-arm-0.28.1.tgz",
|
||||
"integrity": "sha512-0k2F129Xdio1TdJfzJ8sy1Q47vUD2NnwdhiAf7drUN1EBTfPf4hsFCtmMgu/6m8JSzsBrlmVjudMBQqOfG8usQ==",
|
||||
"cpu": [
|
||||
"arm"
|
||||
],
|
||||
|
|
@ -785,9 +785,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/android-arm64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/android-arm64/-/android-arm64-0.27.7.tgz",
|
||||
"integrity": "sha512-62dPZHpIXzvChfvfLJow3q5dDtiNMkwiRzPylSCfriLvZeq0a1bWChrGx/BbUbPwOrsWKMn8idSllklzBy+dgQ==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/android-arm64/-/android-arm64-0.28.1.tgz",
|
||||
"integrity": "sha512-34EGEbCIAgosYz6goLcopX6Mo7NyGv9tfwEM2/7Ce2VcVRk568iSvniGWcUXIy7wEDR1wzolcxcriFVrWYcwBg==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
|
|
@ -802,9 +802,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/android-x64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/android-x64/-/android-x64-0.27.7.tgz",
|
||||
"integrity": "sha512-x5VpMODneVDb70PYV2VQOmIUUiBtY3D3mPBG8NxVk5CogneYhkR7MmM3yR/uMdITLrC1ml/NV1rj4bMJuy9MCg==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/android-x64/-/android-x64-0.28.1.tgz",
|
||||
"integrity": "sha512-dbwY7ltSMDWsRatcRpCnES4F+im88OCUgGZjy52shC7GqHRE/cYlxNbB4Z4UpJswpcc4Qxd2oE/ufM0p61IKng==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
|
|
@ -819,9 +819,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/darwin-arm64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/darwin-arm64/-/darwin-arm64-0.27.7.tgz",
|
||||
"integrity": "sha512-5lckdqeuBPlKUwvoCXIgI2D9/ABmPq3Rdp7IfL70393YgaASt7tbju3Ac+ePVi3KDH6N2RqePfHnXkaDtY9fkw==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/darwin-arm64/-/darwin-arm64-0.28.1.tgz",
|
||||
"integrity": "sha512-TZbWkQY7kvTAXbXUT7uVACR5cMHsDiSz9z7ZKAX/RTq/WJEk3QyRr0wZpNhBDX+/0CtdqUIJlOiodQcta6tY3Q==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
|
|
@ -836,9 +836,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/darwin-x64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/darwin-x64/-/darwin-x64-0.27.7.tgz",
|
||||
"integrity": "sha512-rYnXrKcXuT7Z+WL5K980jVFdvVKhCHhUwid+dDYQpH+qu+TefcomiMAJpIiC2EM3Rjtq0sO3StMV/+3w3MyyqQ==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/darwin-x64/-/darwin-x64-0.28.1.tgz",
|
||||
"integrity": "sha512-zfdzgK9ACBNZLI/CyHTOx81SyNbM6YXn7rxSgX97VjyiPl9W1i4Ka4fgKECEoFCKGpvBj5qArWIGgQjOwkgskQ==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
|
|
@ -853,9 +853,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/freebsd-arm64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/freebsd-arm64/-/freebsd-arm64-0.27.7.tgz",
|
||||
"integrity": "sha512-B48PqeCsEgOtzME2GbNM2roU29AMTuOIN91dsMO30t+Ydis3z/3Ngoj5hhnsOSSwNzS+6JppqWsuhTp6E82l2w==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/freebsd-arm64/-/freebsd-arm64-0.28.1.tgz",
|
||||
"integrity": "sha512-wG2EA8ENdEI0qhkSZMjfqrdY+ziCYCPMmtZjjIwOmXFjmyzEHn+UUxk5of+SYsjtfs3VpnlC7QLzSI5hY/rOAw==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
|
|
@ -870,9 +870,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/freebsd-x64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/freebsd-x64/-/freebsd-x64-0.27.7.tgz",
|
||||
"integrity": "sha512-jOBDK5XEjA4m5IJK3bpAQF9/Lelu/Z9ZcdhTRLf4cajlB+8VEhFFRjWgfy3M1O4rO2GQ/b2dLwCUGpiF/eATNQ==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/freebsd-x64/-/freebsd-x64-0.28.1.tgz",
|
||||
"integrity": "sha512-i7dZ9vQgnvSCzi/rYCXNgtF/U+eKZNJBzu3eTQbRgHnM7tNSizLOkRFAl3qzVc/Op/u5YkHHa4pf/3DOYHthLQ==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
|
|
@ -887,9 +887,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/linux-arm": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/linux-arm/-/linux-arm-0.27.7.tgz",
|
||||
"integrity": "sha512-RkT/YXYBTSULo3+af8Ib0ykH8u2MBh57o7q/DAs3lTJlyVQkgQvlrPTnjIzzRPQyavxtPtfg0EopvDyIt0j1rA==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/linux-arm/-/linux-arm-0.28.1.tgz",
|
||||
"integrity": "sha512-qVXBOHQS+d5Y722GwJzJUtOLlX7km3CraOaGormF1pDtPd2C/l1SHRPgjLunLGe51Sh5YYWKMFDyV4SxgMQYTQ==",
|
||||
"cpu": [
|
||||
"arm"
|
||||
],
|
||||
|
|
@ -904,9 +904,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/linux-arm64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/linux-arm64/-/linux-arm64-0.27.7.tgz",
|
||||
"integrity": "sha512-RZPHBoxXuNnPQO9rvjh5jdkRmVizktkT7TCDkDmQ0W2SwHInKCAV95GRuvdSvA7w4VMwfCjUiPwDi0ZO6Nfe9A==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/linux-arm64/-/linux-arm64-0.28.1.tgz",
|
||||
"integrity": "sha512-yHs+0uc8+nvEAfAfxrWQKK5peSNzBc4PegcMO0EJ2hT71uA7vB8Ihg2e77R2P7SG5uYjPbHlLLmve4LLLRCf0g==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
|
|
@ -921,9 +921,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/linux-ia32": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/linux-ia32/-/linux-ia32-0.27.7.tgz",
|
||||
"integrity": "sha512-GA48aKNkyQDbd3KtkplYWT102C5sn/EZTY4XROkxONgruHPU72l+gW+FfF8tf2cFjeHaRbWpOYa/uRBz/Xq1Pg==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/linux-ia32/-/linux-ia32-0.28.1.tgz",
|
||||
"integrity": "sha512-d1z4ZuP0ajrfz/FhGT4vv278rX8KnPPJx8i5+AtK7TYbx9Le9F1hyzurZpkEyjkGa9dUGhQow4C1NmeGvqxN2w==",
|
||||
"cpu": [
|
||||
"ia32"
|
||||
],
|
||||
|
|
@ -938,9 +938,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/linux-loong64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/linux-loong64/-/linux-loong64-0.27.7.tgz",
|
||||
"integrity": "sha512-a4POruNM2oWsD4WKvBSEKGIiWQF8fZOAsycHOt6JBpZ+JN2n2JH9WAv56SOyu9X5IqAjqSIPTaJkqN8F7XOQ5Q==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/linux-loong64/-/linux-loong64-0.28.1.tgz",
|
||||
"integrity": "sha512-M5sRjUVZrkm1OAPR3dlOYzNmN+loZKGVi1VUQGrwuqLcbR6qeAz+famMhjASeH3YVKvZz+zT1jlh/keC3Rj/lg==",
|
||||
"cpu": [
|
||||
"loong64"
|
||||
],
|
||||
|
|
@ -955,9 +955,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/linux-mips64el": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/linux-mips64el/-/linux-mips64el-0.27.7.tgz",
|
||||
"integrity": "sha512-KabT5I6StirGfIz0FMgl1I+R1H73Gp0ofL9A3nG3i/cYFJzKHhouBV5VWK1CSgKvVaG4q1RNpCTR2LuTVB3fIw==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/linux-mips64el/-/linux-mips64el-0.28.1.tgz",
|
||||
"integrity": "sha512-mRObBZeHh2OxcBFPWE/FjylkRgZdYuiTR3vaTozquCGOH14iP9oN4x4Ge81CoIDYQrXmIxpFumJBu5MtZpnQJQ==",
|
||||
"cpu": [
|
||||
"mips64el"
|
||||
],
|
||||
|
|
@ -972,9 +972,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/linux-ppc64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/linux-ppc64/-/linux-ppc64-0.27.7.tgz",
|
||||
"integrity": "sha512-gRsL4x6wsGHGRqhtI+ifpN/vpOFTQtnbsupUF5R5YTAg+y/lKelYR1hXbnBdzDjGbMYjVJLJTd2OFmMewAgwlQ==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/linux-ppc64/-/linux-ppc64-0.28.1.tgz",
|
||||
"integrity": "sha512-slScBsMAb3GFDcdrCgLwZtPYRoH2H/youv10QiZyRjmsP48fznoveWytSgCI/R0ZcUgpc0ZhIUEx6LHts8yrfQ==",
|
||||
"cpu": [
|
||||
"ppc64"
|
||||
],
|
||||
|
|
@ -989,9 +989,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/linux-riscv64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/linux-riscv64/-/linux-riscv64-0.27.7.tgz",
|
||||
"integrity": "sha512-hL25LbxO1QOngGzu2U5xeXtxXcW+/GvMN3ejANqXkxZ/opySAZMrc+9LY/WyjAan41unrR3YrmtTsUpwT66InQ==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/linux-riscv64/-/linux-riscv64-0.28.1.tgz",
|
||||
"integrity": "sha512-kw0owk1o0GFETUJyW0jc0G4Yzs0BHZn0JDZ8JRT088vjJYX777BAs1fDGxAC+q831qOs2DTC96mNsG2opdfyyQ==",
|
||||
"cpu": [
|
||||
"riscv64"
|
||||
],
|
||||
|
|
@ -1006,9 +1006,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/linux-s390x": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/linux-s390x/-/linux-s390x-0.27.7.tgz",
|
||||
"integrity": "sha512-2k8go8Ycu1Kb46vEelhu1vqEP+UeRVj2zY1pSuPdgvbd5ykAw82Lrro28vXUrRmzEsUV0NzCf54yARIK8r0fdw==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/linux-s390x/-/linux-s390x-0.28.1.tgz",
|
||||
"integrity": "sha512-/lAIjX8aYFRByhh6L5rYtPEDRqa9de/4V/juOXcta5frjvzXO4/sqEtyytse0g3zZFuWu5cDN0MkLz2qRDD2Ag==",
|
||||
"cpu": [
|
||||
"s390x"
|
||||
],
|
||||
|
|
@ -1023,9 +1023,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/linux-x64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/linux-x64/-/linux-x64-0.27.7.tgz",
|
||||
"integrity": "sha512-hzznmADPt+OmsYzw1EE33ccA+HPdIqiCRq7cQeL1Jlq2gb1+OyWBkMCrYGBJ+sxVzve2ZJEVeePbLM2iEIZSxA==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/linux-x64/-/linux-x64-0.28.1.tgz",
|
||||
"integrity": "sha512-u/anNYF2mmVOEDwLtnQ1wOr3EZ9sTNGLWrsYGYwHWzGA3Si84IOkHXlbWTD1NB+9/1lcnweYKO54uhxZydNzfA==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
|
|
@ -1040,9 +1040,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/netbsd-arm64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/netbsd-arm64/-/netbsd-arm64-0.27.7.tgz",
|
||||
"integrity": "sha512-b6pqtrQdigZBwZxAn1UpazEisvwaIDvdbMbmrly7cDTMFnw/+3lVxxCTGOrkPVnsYIosJJXAsILG9XcQS+Yu6w==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/netbsd-arm64/-/netbsd-arm64-0.28.1.tgz",
|
||||
"integrity": "sha512-oks0DYbLwWMmaakTsCb+zL4E+aHRVLom9IJZOAthMQEPiQmydXHkziYEsGYRx0uNV/IjEKGAV941JzH02pflqw==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
|
|
@ -1057,9 +1057,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/netbsd-x64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/netbsd-x64/-/netbsd-x64-0.27.7.tgz",
|
||||
"integrity": "sha512-OfatkLojr6U+WN5EDYuoQhtM+1xco+/6FSzJJnuWiUw5eVcicbyK3dq5EeV/QHT1uy6GoDhGbFpprUiHUYggrw==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/netbsd-x64/-/netbsd-x64-0.28.1.tgz",
|
||||
"integrity": "sha512-aeL6lAnN89Hz43Mlh1G8ARasbuoYvSITDEx0tHh5b7jJnHcssqgjy9Yx430GDpmCa6OyrKoS0aNRjKundRizGg==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
|
|
@ -1074,9 +1074,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/openbsd-arm64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/openbsd-arm64/-/openbsd-arm64-0.27.7.tgz",
|
||||
"integrity": "sha512-AFuojMQTxAz75Fo8idVcqoQWEHIXFRbOc1TrVcFSgCZtQfSdc1RXgB3tjOn/krRHENUB4j00bfGjyl2mJrU37A==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/openbsd-arm64/-/openbsd-arm64-0.28.1.tgz",
|
||||
"integrity": "sha512-MEFJe5C3R8pwXdZ5Y21oo6m7ePiS0d9pWucn99O/wvyJZChoIQKrQDxKrGeW8F5+T0okTHesAmDeiHDTIq0V/Q==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
|
|
@ -1091,9 +1091,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/openbsd-x64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/openbsd-x64/-/openbsd-x64-0.27.7.tgz",
|
||||
"integrity": "sha512-+A1NJmfM8WNDv5CLVQYJ5PshuRm/4cI6WMZRg1by1GwPIQPCTs1GLEUHwiiQGT5zDdyLiRM/l1G0Pv54gvtKIg==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/openbsd-x64/-/openbsd-x64-0.28.1.tgz",
|
||||
"integrity": "sha512-i/ZLIOafE0Z8cI/XANJAixoJL/uRAoS2xOA3rb0xN+KK0K177cMAsQYkzHtBrtMXAKuAc7HGgcWiZ/sRC1Nxgw==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
|
|
@ -1108,9 +1108,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/openharmony-arm64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/openharmony-arm64/-/openharmony-arm64-0.27.7.tgz",
|
||||
"integrity": "sha512-+KrvYb/C8zA9CU/g0sR6w2RBw7IGc5J2BPnc3dYc5VJxHCSF1yNMxTV5LQ7GuKteQXZtspjFbiuW5/dOj7H4Yw==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/openharmony-arm64/-/openharmony-arm64-0.28.1.tgz",
|
||||
"integrity": "sha512-ge+Z7EXFNt2BO1oAMsVpiQ8EwndV9i1xXerAeTIK7AtPs3bKFXQM7nlRxDSIUIMeueR1CNXxqztLzdNeReKBJg==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
|
|
@ -1125,9 +1125,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/sunos-x64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/sunos-x64/-/sunos-x64-0.27.7.tgz",
|
||||
"integrity": "sha512-ikktIhFBzQNt/QDyOL580ti9+5mL/YZeUPKU2ivGtGjdTYoqz6jObj6nOMfhASpS4GU4Q/Clh1QtxWAvcYKamA==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/sunos-x64/-/sunos-x64-0.28.1.tgz",
|
||||
"integrity": "sha512-BEjgtECkL3vY+SaSQ6nzVfiALUeFxpawyp8Jmf5PtYhf1Ug40N1h/hxlhts+f1FvSvarEigdxS3BlSMI2PJLcQ==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
|
|
@ -1142,9 +1142,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/win32-arm64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/win32-arm64/-/win32-arm64-0.27.7.tgz",
|
||||
"integrity": "sha512-7yRhbHvPqSpRUV7Q20VuDwbjW5kIMwTHpptuUzV+AA46kiPze5Z7qgt6CLCK3pWFrHeNfDd1VKgyP4O+ng17CA==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/win32-arm64/-/win32-arm64-0.28.1.tgz",
|
||||
"integrity": "sha512-lCv9eK/H6ZJWbE7bh2nw54CZ9M2nupBxJcTsdk/QQnWkdSjKGuxmmH8/GWrlT1eMmZfn4dGcCjRte397WqfQXA==",
|
||||
"cpu": [
|
||||
"arm64"
|
||||
],
|
||||
|
|
@ -1159,9 +1159,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/win32-ia32": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/win32-ia32/-/win32-ia32-0.27.7.tgz",
|
||||
"integrity": "sha512-SmwKXe6VHIyZYbBLJrhOoCJRB/Z1tckzmgTLfFYOfpMAx63BJEaL9ExI8x7v0oAO3Zh6D/Oi1gVxEYr5oUCFhw==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/win32-ia32/-/win32-ia32-0.28.1.tgz",
|
||||
"integrity": "sha512-zvb/mB2bSCoJOpoCBgYKKpX6YM6mJBlBUVUtVj41DlZJVEB6/0CKlRYxP5wWl1C1ILiCoAU5wZZ4q1P3qeS6Eg==",
|
||||
"cpu": [
|
||||
"ia32"
|
||||
],
|
||||
|
|
@ -1176,9 +1176,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/@esbuild/win32-x64": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/win32-x64/-/win32-x64-0.27.7.tgz",
|
||||
"integrity": "sha512-56hiAJPhwQ1R4i+21FVF7V8kSD5zZTdHcVuRFMW0hn753vVfQN8xlx4uOPT4xoGH0Z/oVATuR82AiqSTDIpaHg==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/@esbuild/win32-x64/-/win32-x64-0.28.1.tgz",
|
||||
"integrity": "sha512-bm4Mowrv+GXMlpWX++EcXw/iLyd1o3+bJkC2DkWXYVvgZCqD/bSj9ctZeAMC3cIxgjRVR2Dufaiu4YPxr5gW1A==",
|
||||
"cpu": [
|
||||
"x64"
|
||||
],
|
||||
|
|
@ -6103,9 +6103,9 @@
|
|||
}
|
||||
},
|
||||
"node_modules/esbuild": {
|
||||
"version": "0.27.7",
|
||||
"resolved": "https://registry.npmjs.org/esbuild/-/esbuild-0.27.7.tgz",
|
||||
"integrity": "sha512-IxpibTjyVnmrIQo5aqNpCgoACA/dTKLTlhMHihVHhdkxKyPO1uBBthumT0rdHmcsk9uMonIWS0m4FljWzILh3w==",
|
||||
"version": "0.28.1",
|
||||
"resolved": "https://registry.npmjs.org/esbuild/-/esbuild-0.28.1.tgz",
|
||||
"integrity": "sha512-HrJrvZv5ayxBzPfwphOoNzkzOIIlifzk0KJrGK2c8R4+LKpMtpYLQeUdjnwjWv/LZlkH2laZk+4w78pi99D4Vw==",
|
||||
"dev": true,
|
||||
"hasInstallScript": true,
|
||||
"license": "MIT",
|
||||
|
|
@ -6116,32 +6116,32 @@
|
|||
"node": ">=18"
|
||||
},
|
||||
"optionalDependencies": {
|
||||
"@esbuild/aix-ppc64": "0.27.7",
|
||||
"@esbuild/android-arm": "0.27.7",
|
||||
"@esbuild/android-arm64": "0.27.7",
|
||||
"@esbuild/android-x64": "0.27.7",
|
||||
"@esbuild/darwin-arm64": "0.27.7",
|
||||
"@esbuild/darwin-x64": "0.27.7",
|
||||
"@esbuild/freebsd-arm64": "0.27.7",
|
||||
"@esbuild/freebsd-x64": "0.27.7",
|
||||
"@esbuild/linux-arm": "0.27.7",
|
||||
"@esbuild/linux-arm64": "0.27.7",
|
||||
"@esbuild/linux-ia32": "0.27.7",
|
||||
"@esbuild/linux-loong64": "0.27.7",
|
||||
"@esbuild/linux-mips64el": "0.27.7",
|
||||
"@esbuild/linux-ppc64": "0.27.7",
|
||||
"@esbuild/linux-riscv64": "0.27.7",
|
||||
"@esbuild/linux-s390x": "0.27.7",
|
||||
"@esbuild/linux-x64": "0.27.7",
|
||||
"@esbuild/netbsd-arm64": "0.27.7",
|
||||
"@esbuild/netbsd-x64": "0.27.7",
|
||||
"@esbuild/openbsd-arm64": "0.27.7",
|
||||
"@esbuild/openbsd-x64": "0.27.7",
|
||||
"@esbuild/openharmony-arm64": "0.27.7",
|
||||
"@esbuild/sunos-x64": "0.27.7",
|
||||
"@esbuild/win32-arm64": "0.27.7",
|
||||
"@esbuild/win32-ia32": "0.27.7",
|
||||
"@esbuild/win32-x64": "0.27.7"
|
||||
"@esbuild/aix-ppc64": "0.28.1",
|
||||
"@esbuild/android-arm": "0.28.1",
|
||||
"@esbuild/android-arm64": "0.28.1",
|
||||
"@esbuild/android-x64": "0.28.1",
|
||||
"@esbuild/darwin-arm64": "0.28.1",
|
||||
"@esbuild/darwin-x64": "0.28.1",
|
||||
"@esbuild/freebsd-arm64": "0.28.1",
|
||||
"@esbuild/freebsd-x64": "0.28.1",
|
||||
"@esbuild/linux-arm": "0.28.1",
|
||||
"@esbuild/linux-arm64": "0.28.1",
|
||||
"@esbuild/linux-ia32": "0.28.1",
|
||||
"@esbuild/linux-loong64": "0.28.1",
|
||||
"@esbuild/linux-mips64el": "0.28.1",
|
||||
"@esbuild/linux-ppc64": "0.28.1",
|
||||
"@esbuild/linux-riscv64": "0.28.1",
|
||||
"@esbuild/linux-s390x": "0.28.1",
|
||||
"@esbuild/linux-x64": "0.28.1",
|
||||
"@esbuild/netbsd-arm64": "0.28.1",
|
||||
"@esbuild/netbsd-x64": "0.28.1",
|
||||
"@esbuild/openbsd-arm64": "0.28.1",
|
||||
"@esbuild/openbsd-x64": "0.28.1",
|
||||
"@esbuild/openharmony-arm64": "0.28.1",
|
||||
"@esbuild/sunos-x64": "0.28.1",
|
||||
"@esbuild/win32-arm64": "0.28.1",
|
||||
"@esbuild/win32-ia32": "0.28.1",
|
||||
"@esbuild/win32-x64": "0.28.1"
|
||||
}
|
||||
},
|
||||
"node_modules/escalade": {
|
||||
|
|
|
|||
|
|
@ -90,7 +90,8 @@
|
|||
"ws": "8.20.1",
|
||||
"braces": "3.0.3",
|
||||
"axios": "1.13.6",
|
||||
"postcss": "8.5.13"
|
||||
"postcss": "8.5.13",
|
||||
"esbuild": "0.28.1"
|
||||
},
|
||||
"engines": {
|
||||
"node": ">=20.9.0",
|
||||
|
|
|
|||
BIN
ui/litellm-dashboard/public/assets/logos/cisco.png
Normal file
BIN
ui/litellm-dashboard/public/assets/logos/cisco.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 1.9 KiB |
|
|
@ -8,6 +8,7 @@ import {
|
|||
useModelHub,
|
||||
useModelsInfo,
|
||||
useSelectedTeamModels,
|
||||
useUserModels,
|
||||
type AllProxyModelsResponse,
|
||||
type PaginatedModelInfoResponse,
|
||||
type ProxyModel,
|
||||
|
|
@ -480,6 +481,70 @@ describe("useAllProxyModels", () => {
|
|||
});
|
||||
});
|
||||
|
||||
describe("useUserModels", () => {
|
||||
let queryClient: QueryClient;
|
||||
|
||||
beforeEach(() => {
|
||||
queryClient = new QueryClient({
|
||||
defaultOptions: {
|
||||
queries: {
|
||||
retry: false,
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
vi.clearAllMocks();
|
||||
|
||||
mockUseAuthorized.mockReturnValue({
|
||||
accessToken: "test-access-token",
|
||||
userId: "test-user-id",
|
||||
userRole: "Admin",
|
||||
token: "test-token",
|
||||
userEmail: "test@example.com",
|
||||
premiumUser: false,
|
||||
disabledPersonalKeyCreation: null,
|
||||
showSSOBanner: false,
|
||||
});
|
||||
});
|
||||
|
||||
const wrapper = ({ children }: { children: ReactNode }) =>
|
||||
React.createElement(QueryClientProvider, { client: queryClient }, children);
|
||||
|
||||
it("maps the available-models response to a list of model ids", async () => {
|
||||
(modelAvailableCall as any).mockResolvedValue({
|
||||
data: [{ id: "gpt-4" }, { id: "claude-3-opus" }],
|
||||
});
|
||||
|
||||
const { result } = renderHook(() => useUserModels(), { wrapper });
|
||||
|
||||
await waitFor(() => {
|
||||
expect(result.current.isSuccess).toBe(true);
|
||||
});
|
||||
|
||||
expect(result.current.data).toEqual(["gpt-4", "claude-3-opus"]);
|
||||
expect(modelAvailableCall).toHaveBeenCalledWith("test-access-token", "test-user-id", "Admin");
|
||||
expect(modelAvailableCall).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
it("should not execute query when accessToken is missing", () => {
|
||||
mockUseAuthorized.mockReturnValue({
|
||||
accessToken: null,
|
||||
userId: "test-user-id",
|
||||
userRole: "Admin",
|
||||
token: null,
|
||||
userEmail: "test@example.com",
|
||||
premiumUser: false,
|
||||
disabledPersonalKeyCreation: null,
|
||||
showSSOBanner: false,
|
||||
});
|
||||
|
||||
const { result } = renderHook(() => useUserModels(), { wrapper });
|
||||
|
||||
expect(result.current.isFetched).toBe(false);
|
||||
expect(modelAvailableCall).not.toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
||||
describe("useSelectedTeamModels", () => {
|
||||
let queryClient: QueryClient;
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
import { useQuery, useInfiniteQuery } from "@tanstack/react-query";
|
||||
import { useQuery, useInfiniteQuery, UseQueryResult } from "@tanstack/react-query";
|
||||
import { createQueryKeys } from "../common/queryKeysFactory";
|
||||
import { modelInfoCall, modelHubCall, modelAvailableCall } from "@/components/networking";
|
||||
import useAuthorized from "../useAuthorized";
|
||||
|
|
@ -27,6 +27,7 @@ const modelHubKeys = createQueryKeys("modelHub");
|
|||
const allProxyModelsKeys = createQueryKeys("allProxyModels");
|
||||
const selectedTeamModelsKeys = createQueryKeys("selectedTeamModels");
|
||||
const infiniteModelKeys = createQueryKeys("infiniteModels");
|
||||
const userModelsKeys = createQueryKeys("userModels");
|
||||
|
||||
export const useModelsInfo = (
|
||||
page: number = 1,
|
||||
|
|
@ -76,6 +77,18 @@ export const useAllProxyModels = () => {
|
|||
});
|
||||
};
|
||||
|
||||
export const useUserModels = (): UseQueryResult<string[]> => {
|
||||
const { accessToken, userId, userRole } = useAuthorized();
|
||||
return useQuery<string[]>({
|
||||
queryKey: userModelsKeys.list({}),
|
||||
queryFn: async () => {
|
||||
const response = await modelAvailableCall(accessToken!, userId!, userRole!);
|
||||
return response["data"].map((model: { id: string }) => model.id);
|
||||
},
|
||||
enabled: Boolean(accessToken && userId && userRole),
|
||||
});
|
||||
};
|
||||
|
||||
export const useSelectedTeamModels = (teamID: string | null) => {
|
||||
const { accessToken, userId, userRole } = useAuthorized();
|
||||
return useQuery<AllProxyModelsResponse>({
|
||||
|
|
|
|||
|
|
@ -2,13 +2,14 @@ import { describe, it, expect, vi, beforeEach } from "vitest";
|
|||
import { renderHook, waitFor } from "@testing-library/react";
|
||||
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
|
||||
import React, { ReactNode } from "react";
|
||||
import { useOrganizations } from "./useOrganizations";
|
||||
import { organizationListCall } from "@/components/networking";
|
||||
import { organizationKeys, useOrganization, useOrganizations } from "./useOrganizations";
|
||||
import { organizationInfoCall, organizationListCall } from "@/components/networking";
|
||||
import type { Organization } from "@/components/networking";
|
||||
|
||||
// Mock the networking function
|
||||
vi.mock("@/components/networking", () => ({
|
||||
organizationListCall: vi.fn(),
|
||||
organizationInfoCall: vi.fn(),
|
||||
}));
|
||||
|
||||
// Mock useAuthorized hook - we can override this in individual tests
|
||||
|
|
@ -107,7 +108,7 @@ describe("useOrganizations", () => {
|
|||
|
||||
expect(result.current.data).toEqual(mockOrganizations);
|
||||
expect(result.current.error).toBeNull();
|
||||
expect(organizationListCall).toHaveBeenCalledWith("test-access-token");
|
||||
expect(organizationListCall).toHaveBeenCalledWith("test-access-token", null, null);
|
||||
expect(organizationListCall).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
|
|
@ -131,10 +132,47 @@ describe("useOrganizations", () => {
|
|||
|
||||
expect(result.current.error).toEqual(testError);
|
||||
expect(result.current.data).toBeUndefined();
|
||||
expect(organizationListCall).toHaveBeenCalledWith("test-access-token");
|
||||
expect(organizationListCall).toHaveBeenCalledWith("test-access-token", null, null);
|
||||
expect(organizationListCall).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
it("passes org_id and org_alias filters to organizationListCall and caches separately from the unfiltered list", async () => {
|
||||
(organizationListCall as any).mockResolvedValue(mockOrganizations);
|
||||
|
||||
const { result } = renderHook(() => useOrganizations({ org_id: "org-1", org_alias: "Test Organization 1" }), {
|
||||
wrapper,
|
||||
});
|
||||
|
||||
await waitFor(() => {
|
||||
expect(result.current.isSuccess).toBe(true);
|
||||
});
|
||||
|
||||
expect(organizationListCall).toHaveBeenCalledWith("test-access-token", "org-1", "Test Organization 1");
|
||||
|
||||
(organizationListCall as any).mockResolvedValue([]);
|
||||
const { result: unfiltered } = renderHook(() => useOrganizations(), { wrapper });
|
||||
|
||||
await waitFor(() => {
|
||||
expect(unfiltered.current.isSuccess).toBe(true);
|
||||
});
|
||||
|
||||
expect(organizationListCall).toHaveBeenLastCalledWith("test-access-token", null, null);
|
||||
expect(organizationListCall).toHaveBeenCalledTimes(2);
|
||||
});
|
||||
|
||||
it("treats empty-string filters as no filters, writing to the unfiltered cache entry", async () => {
|
||||
(organizationListCall as any).mockResolvedValue(mockOrganizations);
|
||||
|
||||
const { result } = renderHook(() => useOrganizations({ org_id: "", org_alias: "" }), { wrapper });
|
||||
|
||||
await waitFor(() => {
|
||||
expect(result.current.isSuccess).toBe(true);
|
||||
});
|
||||
|
||||
expect(organizationListCall).toHaveBeenCalledWith("test-access-token", null, null);
|
||||
expect(queryClient.getQueryData(organizationKeys.list({}))).toEqual(mockOrganizations);
|
||||
});
|
||||
|
||||
it("should not execute query when accessToken is missing", async () => {
|
||||
// Mock missing accessToken
|
||||
mockUseAuthorized.mockReturnValue({
|
||||
|
|
@ -243,7 +281,7 @@ describe("useOrganizations", () => {
|
|||
expect(result.current.isLoading).toBe(false);
|
||||
});
|
||||
|
||||
expect(organizationListCall).toHaveBeenCalledWith("test-access-token");
|
||||
expect(organizationListCall).toHaveBeenCalledWith("test-access-token", null, null);
|
||||
expect(organizationListCall).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
|
|
@ -260,7 +298,7 @@ describe("useOrganizations", () => {
|
|||
});
|
||||
|
||||
expect(result.current.data).toEqual([]);
|
||||
expect(organizationListCall).toHaveBeenCalledWith("test-access-token");
|
||||
expect(organizationListCall).toHaveBeenCalledWith("test-access-token", null, null);
|
||||
});
|
||||
|
||||
it("should handle network timeout error", async () => {
|
||||
|
|
@ -280,3 +318,58 @@ describe("useOrganizations", () => {
|
|||
expect(result.current.data).toBeUndefined();
|
||||
});
|
||||
});
|
||||
|
||||
describe("useOrganization", () => {
|
||||
let queryClient: QueryClient;
|
||||
|
||||
beforeEach(() => {
|
||||
queryClient = new QueryClient({
|
||||
defaultOptions: { queries: { retry: false } },
|
||||
});
|
||||
|
||||
vi.clearAllMocks();
|
||||
|
||||
mockUseAuthorized.mockReturnValue({
|
||||
accessToken: "test-access-token",
|
||||
userId: "test-user-id",
|
||||
userRole: "Admin",
|
||||
token: "test-token",
|
||||
userEmail: "test@example.com",
|
||||
premiumUser: false,
|
||||
disabledPersonalKeyCreation: null,
|
||||
showSSOBanner: false,
|
||||
});
|
||||
});
|
||||
|
||||
const wrapper = ({ children }: { children: ReactNode }) =>
|
||||
React.createElement(QueryClientProvider, { client: queryClient }, children);
|
||||
|
||||
it("seeds initialData from a filtered list cache entry so the detail renders without a loading state", () => {
|
||||
(organizationInfoCall as any).mockResolvedValue(mockOrganizations[1]);
|
||||
// Only a filtered list was ever fetched; the unfiltered list({}) entry stays empty.
|
||||
queryClient.setQueryData(organizationKeys.list({ filters: { org_id: "org-2" } }), [mockOrganizations[1]]);
|
||||
|
||||
const { result } = renderHook(() => useOrganization("org-2"), { wrapper });
|
||||
|
||||
// initialData found org-2 in the filtered cache, so data is present on the first render.
|
||||
expect(result.current.data).toEqual(mockOrganizations[1]);
|
||||
expect(result.current.isLoading).toBe(false);
|
||||
});
|
||||
|
||||
it("falls through to the detail API call when no cached list contains the organization", async () => {
|
||||
(organizationInfoCall as any).mockResolvedValue(mockOrganizations[0]);
|
||||
queryClient.setQueryData(organizationKeys.list({ filters: { org_id: "org-2" } }), [mockOrganizations[1]]);
|
||||
|
||||
const { result } = renderHook(() => useOrganization("org-1"), { wrapper });
|
||||
|
||||
// org-1 is in no cached list, so there is no initialData and it loads via the detail API.
|
||||
expect(result.current.data).toBeUndefined();
|
||||
|
||||
await waitFor(() => {
|
||||
expect(result.current.isSuccess).toBe(true);
|
||||
});
|
||||
|
||||
expect(organizationInfoCall).toHaveBeenCalledWith("test-access-token", "org-1");
|
||||
expect(result.current.data).toEqual(mockOrganizations[0]);
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -4,11 +4,23 @@ import { useQuery, useQueryClient, UseQueryResult } from "@tanstack/react-query"
|
|||
import { createQueryKeys } from "../common/queryKeysFactory";
|
||||
|
||||
export const organizationKeys = createQueryKeys("organizations");
|
||||
export const useOrganizations = (): UseQueryResult<Organization[]> => {
|
||||
|
||||
export interface OrganizationListFilters {
|
||||
org_id?: string | null;
|
||||
org_alias?: string | null;
|
||||
}
|
||||
|
||||
export const useOrganizations = (filters?: OrganizationListFilters): UseQueryResult<Organization[]> => {
|
||||
const { accessToken, userId, userRole } = useAuthorized();
|
||||
const orgId = filters?.org_id || null;
|
||||
const orgAlias = filters?.org_alias || null;
|
||||
return useQuery<Organization[]>({
|
||||
queryKey: organizationKeys.list({}),
|
||||
queryFn: async () => await organizationListCall(accessToken!),
|
||||
queryKey: organizationKeys.list(
|
||||
orgId || orgAlias
|
||||
? { filters: { ...(orgId && { org_id: orgId }), ...(orgAlias && { org_alias: orgAlias }) } }
|
||||
: {},
|
||||
),
|
||||
queryFn: async () => await organizationListCall(accessToken!, orgId, orgAlias),
|
||||
enabled: Boolean(accessToken && userId && userRole),
|
||||
});
|
||||
};
|
||||
|
|
@ -31,9 +43,10 @@ export const useOrganization = (organizationID?: string) => {
|
|||
initialData: () => {
|
||||
if (!organizationID) return undefined;
|
||||
|
||||
const organizations = queryClient.getQueryData<Organization[]>(organizationKeys.list({}));
|
||||
|
||||
return organizations?.find((organization: Organization) => organization.organization_id === organizationID);
|
||||
return queryClient
|
||||
.getQueriesData<Organization[]>({ queryKey: organizationKeys.lists() })
|
||||
.flatMap(([, organizations]) => organizations ?? [])
|
||||
.find((organization) => organization.organization_id === organizationID);
|
||||
},
|
||||
});
|
||||
};
|
||||
|
|
|
|||
|
|
@ -0,0 +1,9 @@
|
|||
"use client";
|
||||
|
||||
import OrganizationsTable from "@/components/organizations";
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
|
||||
export default function OrganizationsPage() {
|
||||
const { accessToken, userRole, premiumUser } = useAuthorized();
|
||||
return <OrganizationsTable userRole={userRole ?? ""} accessToken={accessToken} premiumUser={premiumUser ?? false} />;
|
||||
}
|
||||
|
|
@ -6,8 +6,8 @@ import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings"
|
|||
import LoadingScreen from "@/components/common_components/LoadingScreen";
|
||||
import { Team } from "@/components/key_team_helpers/key_list";
|
||||
import { Organization, proxyBaseUrl, getInProductNudgesCall } from "@/components/networking";
|
||||
import { fetchUserModels, CreateKeyPrefillData } from "@/components/organisms/create_key_button";
|
||||
import Organizations, { fetchOrganizations } from "@/components/organizations";
|
||||
import { CreateKeyPrefillData } from "@/components/organisms/create_key_button";
|
||||
import { fetchOrganizations } from "@/components/organizations";
|
||||
import PassThroughSettings from "@/components/pass_through_settings";
|
||||
import { SurveyPrompt, SurveyModal, ClaudeCodePrompt, ClaudeCodeModal } from "@/components/survey";
|
||||
import Usage from "@/components/usage";
|
||||
|
|
@ -31,7 +31,6 @@ function CreateKeyPageContent() {
|
|||
const [teams, setTeams] = useState<Team[] | null>(null);
|
||||
const [keys, setKeys] = useState<null | any[]>([]);
|
||||
const [organizations, setOrganizations] = useState<Organization[]>([]);
|
||||
const [userModels, setUserModels] = useState<string[]>([]);
|
||||
|
||||
const router = useRouter();
|
||||
const searchParams = useSearchParams()!;
|
||||
|
|
@ -171,9 +170,6 @@ function CreateKeyPageContent() {
|
|||
}, [token]);
|
||||
|
||||
useEffect(() => {
|
||||
if (accessToken && userID && userRole) {
|
||||
fetchUserModels(userID, userRole, accessToken, setUserModels);
|
||||
}
|
||||
if (accessToken && userID && userRole) {
|
||||
v2TeamListCall(accessToken, 1, 100, {
|
||||
userID: userRole !== "Admin" && userRole !== "Admin Viewer" ? userID : null,
|
||||
|
|
@ -321,15 +317,6 @@ function CreateKeyPageContent() {
|
|||
premiumUser={premiumUser}
|
||||
teams={teams}
|
||||
/>
|
||||
) : page == "organizations" ? (
|
||||
<Organizations
|
||||
organizations={organizations}
|
||||
setOrganizations={setOrganizations}
|
||||
userModels={userModels}
|
||||
accessToken={accessToken}
|
||||
userRole={userRole}
|
||||
premiumUser={premiumUser}
|
||||
/>
|
||||
) : page == "pass-through-settings" ? (
|
||||
<PassThroughSettings
|
||||
userID={userID}
|
||||
|
|
|
|||
|
|
@ -1,648 +0,0 @@
|
|||
import { useState, useRef, useEffect, useCallback } from "react";
|
||||
|
||||
const STYLES = `
|
||||
.wrt-wrap {
|
||||
font-family: 'JetBrains Mono', 'Fira Code', monospace;
|
||||
background: #0d0d14;
|
||||
border: 1px solid #1e1e2e;
|
||||
border-radius: 10px;
|
||||
overflow: hidden;
|
||||
margin: 24px 0;
|
||||
}
|
||||
|
||||
.wrt-toggle {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
padding: 14px 20px;
|
||||
cursor: pointer;
|
||||
user-select: none;
|
||||
background: #0d0d14;
|
||||
transition: background 0.15s;
|
||||
}
|
||||
.wrt-toggle:hover { background: #111120; }
|
||||
|
||||
.wrt-toggle-left { display: flex; align-items: center; gap: 10px; }
|
||||
|
||||
.wrt-live-dot {
|
||||
width: 8px; height: 8px; border-radius: 50%;
|
||||
background: #00ff88;
|
||||
box-shadow: 0 0 8px #00ff88;
|
||||
animation: wrt-blink 2s infinite;
|
||||
}
|
||||
@keyframes wrt-blink { 0%,100%{opacity:1} 50%{opacity:0.4} }
|
||||
|
||||
.wrt-toggle-title { font-size: 12px; font-weight: 600; color: #e2e8f0; letter-spacing: 0.06em; }
|
||||
.wrt-toggle-sub { font-size: 10px; color: #4a5568; margin-top: 1px; }
|
||||
.wrt-chevron { font-size: 11px; color: #4a5568; transition: transform 0.2s; }
|
||||
.wrt-chevron.open { transform: rotate(180deg); }
|
||||
|
||||
.wrt-body {
|
||||
border-top: 1px solid #1e1e2e;
|
||||
display: grid;
|
||||
grid-template-columns: 280px 1fr;
|
||||
height: 460px;
|
||||
}
|
||||
|
||||
.wrt-sidebar {
|
||||
border-right: 1px solid #1e1e2e;
|
||||
padding: 14px;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 12px;
|
||||
overflow-y: auto;
|
||||
}
|
||||
|
||||
.wrt-label {
|
||||
font-size: 9px;
|
||||
letter-spacing: 0.15em;
|
||||
color: #4a5568;
|
||||
text-transform: uppercase;
|
||||
margin-bottom: 5px;
|
||||
}
|
||||
|
||||
.wrt-field { display: flex; flex-direction: column; gap: 4px; margin-bottom: 6px; }
|
||||
.wrt-field label { font-size: 10px; color: #4a5568; }
|
||||
.wrt-field input {
|
||||
background: #0a0a0f;
|
||||
border: 1px solid #1e1e2e;
|
||||
border-radius: 5px;
|
||||
color: #e2e8f0;
|
||||
font-family: inherit;
|
||||
font-size: 11px;
|
||||
padding: 7px 9px;
|
||||
outline: none;
|
||||
width: 100%;
|
||||
transition: border-color 0.2s;
|
||||
}
|
||||
.wrt-field input:focus { border-color: #7c3aed; }
|
||||
|
||||
.wrt-divider { height: 1px; background: #1e1e2e; }
|
||||
|
||||
.wrt-btn {
|
||||
display: flex; align-items: center; justify-content: center;
|
||||
border: none; border-radius: 5px; cursor: pointer;
|
||||
font-family: inherit; font-size: 11px; font-weight: 600;
|
||||
padding: 8px; width: 100%;
|
||||
transition: all 0.15s; letter-spacing: 0.04em;
|
||||
}
|
||||
.wrt-btn + .wrt-btn { margin-top: 5px; }
|
||||
.wrt-btn-primary { background: #00ff88; color: #000; }
|
||||
.wrt-btn-primary:hover:not(:disabled) { filter: brightness(1.1); }
|
||||
.wrt-btn-primary:disabled { opacity: 0.35; cursor: not-allowed; }
|
||||
.wrt-btn-danger { background: transparent; color: #ff4466; border: 1px solid #ff4466; }
|
||||
.wrt-btn-danger:hover:not(:disabled) { background: rgba(255,68,102,0.08); }
|
||||
.wrt-btn-danger:disabled { opacity: 0.3; cursor: not-allowed; }
|
||||
.wrt-btn-ghost { background: #111118; color: #e2e8f0; border: 1px solid #1e1e2e; }
|
||||
.wrt-btn-ghost:hover { border-color: #7c3aed; }
|
||||
|
||||
.wrt-flow { display: flex; align-items: center; padding: 4px 0; gap: 0; }
|
||||
.wrt-flow-box {
|
||||
padding: 4px 7px; border-radius: 4px; font-size: 9px;
|
||||
border: 1px solid #1e1e2e; color: #4a5568;
|
||||
transition: all 0.3s; white-space: nowrap;
|
||||
}
|
||||
.wrt-flow-box.active { border-color: #00ff88; color: #00ff88; box-shadow: 0 0 8px rgba(0,255,136,0.15); }
|
||||
.wrt-flow-arrow { font-size: 10px; color: #4a5568; padding: 0 4px; transition: color 0.3s; }
|
||||
.wrt-flow-arrow.active { color: #00ff88; }
|
||||
|
||||
.wrt-meta { display: flex; flex-direction: column; gap: 4px; }
|
||||
.wrt-meta-row { display: flex; justify-content: space-between; font-size: 10px; }
|
||||
.wrt-meta-row span:first-child { color: #4a5568; }
|
||||
.wrt-meta-row span:last-child { color: #e2e8f0; }
|
||||
|
||||
.wrt-status-pill {
|
||||
display: flex; align-items: center; gap: 6px;
|
||||
font-size: 10px; color: #4a5568;
|
||||
background: #111118; border: 1px solid #1e1e2e;
|
||||
border-radius: 100px; padding: 3px 10px;
|
||||
}
|
||||
.wrt-status-dot {
|
||||
width: 6px; height: 6px; border-radius: 50%;
|
||||
background: #4a5568; transition: all 0.3s;
|
||||
}
|
||||
.wrt-status-dot.connected { background: #00ff88; box-shadow: 0 0 6px #00ff88; }
|
||||
.wrt-status-dot.connecting { background: #ffaa00; animation: wrt-blink 1s infinite; }
|
||||
.wrt-status-dot.error { background: #ff4466; }
|
||||
|
||||
.wrt-main { display: flex; flex-direction: column; overflow: hidden; }
|
||||
|
||||
.wrt-header {
|
||||
display: flex; align-items: center; justify-content: space-between;
|
||||
padding: 8px 14px; border-bottom: 1px solid #1e1e2e; background: #111118;
|
||||
}
|
||||
.wrt-header-title { font-size: 10px; color: #4a5568; letter-spacing: 0.08em; }
|
||||
|
||||
.wrt-tabs { display: flex; padding: 0 14px; border-bottom: 1px solid #1e1e2e; }
|
||||
.wrt-tab {
|
||||
font-size: 9px; letter-spacing: 0.08em; padding: 10px 12px; cursor: pointer;
|
||||
color: #4a5568; border-bottom: 2px solid transparent; transition: all 0.15s;
|
||||
user-select: none;
|
||||
}
|
||||
.wrt-tab.active { color: #00ff88; border-bottom-color: #00ff88; }
|
||||
.wrt-tab:hover:not(.active) { color: #e2e8f0; }
|
||||
|
||||
.wrt-tab-content { flex: 1; overflow: hidden; display: none; flex-direction: column; }
|
||||
.wrt-tab-content.active { display: flex; }
|
||||
|
||||
.wrt-log {
|
||||
flex: 1; overflow-y: auto; padding: 8px 12px;
|
||||
display: flex; flex-direction: column; gap: 2px;
|
||||
}
|
||||
.wrt-log::-webkit-scrollbar { width: 3px; }
|
||||
.wrt-log::-webkit-scrollbar-thumb { background: #1e1e2e; border-radius: 2px; }
|
||||
|
||||
.wrt-entry {
|
||||
display: grid; grid-template-columns: 58px 56px 1fr; gap: 8px;
|
||||
padding: 3px 7px; border-radius: 3px;
|
||||
border-left: 2px solid transparent;
|
||||
font-size: 10px; line-height: 1.5;
|
||||
animation: wrt-fadein 0.15s ease;
|
||||
}
|
||||
@keyframes wrt-fadein { from { opacity:0; transform:translateY(2px); } to { opacity:1; transform:none; } }
|
||||
|
||||
.wrt-entry.info { border-left-color: #7c3aed; }
|
||||
.wrt-entry.info .we-tag { color: #7c3aed; }
|
||||
.wrt-entry.success { border-left-color: #00ff88; }
|
||||
.wrt-entry.success .we-tag { color: #00ff88; }
|
||||
.wrt-entry.error { border-left-color: #ff4466; }
|
||||
.wrt-entry.error .we-tag { color: #ff4466; }
|
||||
.wrt-entry.warn { border-left-color: #ffaa00; }
|
||||
.wrt-entry.warn .we-tag { color: #ffaa00; }
|
||||
.wrt-entry.step { border-left-color: #60a5fa; }
|
||||
.wrt-entry.step .we-tag { color: #60a5fa; }
|
||||
|
||||
.we-time { color: #4a5568; font-size: 9px; padding-top: 1px; }
|
||||
.we-tag { font-size: 9px; font-weight: 700; padding-top: 1px; }
|
||||
.we-msg { color: #e2e8f0; word-break: break-all; white-space: pre-wrap; }
|
||||
|
||||
.wrt-empty {
|
||||
display: flex; flex-direction: column; align-items: center; justify-content: center;
|
||||
flex: 1; gap: 6px; color: #4a5568; font-size: 11px;
|
||||
}
|
||||
|
||||
.wrt-sdp-pane { flex: 1; display: grid; grid-template-columns: 1fr 1fr; overflow: hidden; }
|
||||
.wrt-sdp-box { display: flex; flex-direction: column; border-right: 1px solid #1e1e2e; overflow: hidden; }
|
||||
.wrt-sdp-box:last-child { border-right: none; }
|
||||
.wrt-sdp-hdr {
|
||||
padding: 7px 12px; border-bottom: 1px solid #1e1e2e;
|
||||
font-size: 9px; color: #4a5568; letter-spacing: 0.08em;
|
||||
display: flex; align-items: center; gap: 6px;
|
||||
}
|
||||
.wrt-sdp-dot { width: 5px; height: 5px; border-radius: 50%; background: #1e1e2e; }
|
||||
.wrt-sdp-dot.active { background: #00ff88; }
|
||||
.wrt-sdp-pane textarea {
|
||||
flex: 1; background: transparent; border: none; color: #e2e8f0;
|
||||
font-family: inherit; font-size: 10px; padding: 10px 12px;
|
||||
resize: none; outline: none; line-height: 1.5;
|
||||
}
|
||||
|
||||
.wrt-audio-pane {
|
||||
flex: 1; display: flex; flex-direction: column;
|
||||
align-items: center; justify-content: center; gap: 14px;
|
||||
}
|
||||
.wrt-viz { display: flex; align-items: center; gap: 2px; height: 44px; }
|
||||
.wrt-bar { width: 3px; border-radius: 2px; min-height: 2px; background: #00ff88; transition: height 0.05s; }
|
||||
.wrt-mic-btn {
|
||||
width: 52px; height: 52px; border-radius: 50%;
|
||||
background: #111118; border: 1.5px solid #1e1e2e;
|
||||
font-size: 18px; cursor: pointer;
|
||||
display: flex; align-items: center; justify-content: center; transition: all 0.2s;
|
||||
}
|
||||
.wrt-mic-btn.active { border-color: #00ff88; box-shadow: 0 0 16px rgba(0,255,136,0.2); }
|
||||
.wrt-audio-status { font-size: 10px; color: #4a5568; text-align: center; }
|
||||
`;
|
||||
|
||||
function useLog() {
|
||||
const [entries, setEntries] = useState([]);
|
||||
const add = useCallback((level, tag, msg) => {
|
||||
const time = new Date().toTimeString().slice(0, 8);
|
||||
setEntries((prev) => [...prev, { level, tag, msg, time, id: Date.now() + Math.random() }]);
|
||||
}, []);
|
||||
const clear = useCallback(() => setEntries([]), []);
|
||||
return { entries, add, clear };
|
||||
}
|
||||
|
||||
export default function WebRTCTester() {
|
||||
const [open, setOpen] = useState(false);
|
||||
const [activeTab, setActiveTab] = useState("logs");
|
||||
const [proxyUrl, setProxyUrl] = useState("http://localhost:4000");
|
||||
const [apiKey, setApiKey] = useState("sk-1234");
|
||||
const [model, setModel] = useState("gpt-4o-realtime-preview");
|
||||
const [status, setStatus] = useState("idle");
|
||||
const [flowStep, setFlowStep] = useState(0);
|
||||
const [tokenPreview, setTokenPreview] = useState("—");
|
||||
const [iceState, setIceState] = useState("—");
|
||||
const [connState, setConnState] = useState("—");
|
||||
const [dcState, setDcState] = useState("—");
|
||||
const [sdpOffer, setSdpOffer] = useState("");
|
||||
const [sdpAnswer, setSdpAnswer] = useState("");
|
||||
const [offerActive, setOfferActive] = useState(false);
|
||||
const [answerActive, setAnswerActive] = useState(false);
|
||||
const [audioStatus, setAudioStatus] = useState("Start a session first");
|
||||
const [micActive, setMicActive] = useState(false);
|
||||
const [bars, setBars] = useState(Array(28).fill(2));
|
||||
const [connected, setConnected] = useState(false);
|
||||
|
||||
const { entries, add: log, clear: clearLogs } = useLog();
|
||||
const logRef = useRef(null);
|
||||
|
||||
const pcRef = useRef(null);
|
||||
const dcRef = useRef(null);
|
||||
const streamRef = useRef(null);
|
||||
const audioCtxRef = useRef(null);
|
||||
const analyserRef = useRef(null);
|
||||
const animRef = useRef(null);
|
||||
const tokenRef = useRef(null);
|
||||
const micRef = useRef(false);
|
||||
const remoteAudioRef = useRef(null);
|
||||
|
||||
useEffect(() => {
|
||||
if (logRef.current) logRef.current.scrollTop = logRef.current.scrollHeight;
|
||||
}, [entries]);
|
||||
|
||||
function drawBars() {
|
||||
animRef.current = requestAnimationFrame(drawBars);
|
||||
if (!analyserRef.current) return;
|
||||
const data = new Uint8Array(analyserRef.current.frequencyBinCount);
|
||||
analyserRef.current.getByteFrequencyData(data);
|
||||
setBars(Array.from({ length: 28 }, (_, i) => Math.max(2, ((data[i] || 0) / 255) * 42)));
|
||||
}
|
||||
|
||||
function setupAnalyser(stream) {
|
||||
audioCtxRef.current = new AudioContext();
|
||||
const src = audioCtxRef.current.createMediaStreamSource(stream);
|
||||
analyserRef.current = audioCtxRef.current.createAnalyser();
|
||||
analyserRef.current.fftSize = 64;
|
||||
src.connect(analyserRef.current);
|
||||
drawBars();
|
||||
}
|
||||
|
||||
async function startSession() {
|
||||
const url = proxyUrl.trim().replace(/\/$/, "");
|
||||
const key = apiKey.trim();
|
||||
const mdl = model.trim();
|
||||
|
||||
setConnected(true);
|
||||
setStatus("connecting");
|
||||
setFlowStep(1);
|
||||
|
||||
// Step 1: ephemeral token
|
||||
log("step", "STEP 1", `POST ${url}/v1/realtime/client_secrets`);
|
||||
let tokenResp;
|
||||
try {
|
||||
const r = await fetch(`${url}/v1/realtime/client_secrets`, {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json", Authorization: `Bearer ${key}` },
|
||||
body: JSON.stringify({ model: mdl }),
|
||||
});
|
||||
log("info", "HTTP", `${r.status} ${r.statusText}`);
|
||||
const raw = await r.text();
|
||||
if (!r.ok) {
|
||||
log("error", "ERR", raw);
|
||||
stopSession();
|
||||
return;
|
||||
}
|
||||
tokenResp = JSON.parse(raw);
|
||||
log("success", "TOKEN", "Received encrypted ephemeral token");
|
||||
} catch (e) {
|
||||
log("error", "ERR", `client_secrets failed: ${e.message}`);
|
||||
stopSession();
|
||||
return;
|
||||
}
|
||||
|
||||
const token = tokenResp?.client_secret?.value ?? tokenResp?.value;
|
||||
if (!token) {
|
||||
log("error", "ERR", `Cannot extract token: ${JSON.stringify(tokenResp)}`);
|
||||
stopSession();
|
||||
return;
|
||||
}
|
||||
tokenRef.current = token;
|
||||
setTokenPreview(token.slice(0, 10) + "…");
|
||||
log("info", "TOKEN", `Preview: ${token.slice(0, 10)}…`);
|
||||
|
||||
// Step 2: PeerConnection
|
||||
log("step", "STEP 2", "Creating RTCPeerConnection");
|
||||
const pc = new RTCPeerConnection();
|
||||
pcRef.current = pc;
|
||||
|
||||
pc.oniceconnectionstatechange = () => {
|
||||
setIceState(pc.iceConnectionState);
|
||||
log("info", "ICE", pc.iceConnectionState);
|
||||
if (pc.iceConnectionState === "connected" || pc.iceConnectionState === "completed") {
|
||||
setStatus("connected");
|
||||
setFlowStep(3);
|
||||
}
|
||||
if (pc.iceConnectionState === "failed" || pc.iceConnectionState === "disconnected") {
|
||||
setStatus("error");
|
||||
}
|
||||
};
|
||||
|
||||
pc.onconnectionstatechange = () => {
|
||||
setConnState(pc.connectionState);
|
||||
log("info", "CONN", pc.connectionState);
|
||||
};
|
||||
|
||||
pc.ontrack = (e) => {
|
||||
log("success", "AUDIO", "Remote audio track received from OpenAI");
|
||||
if (remoteAudioRef.current) remoteAudioRef.current.srcObject = e.streams[0];
|
||||
setupAnalyser(e.streams[0]);
|
||||
setAudioStatus("Receiving audio from OpenAI ✓");
|
||||
};
|
||||
|
||||
const dc = pc.createDataChannel("oai-events");
|
||||
dcRef.current = dc;
|
||||
dc.onopen = () => {
|
||||
setDcState("open");
|
||||
log("success", "DC", "Data channel open — ready!");
|
||||
setStatus("connected");
|
||||
};
|
||||
dc.onclose = () => {
|
||||
setDcState("closed");
|
||||
log("warn", "DC", "Closed");
|
||||
};
|
||||
dc.onmessage = (e) => {
|
||||
try {
|
||||
log("info", "EVENT", JSON.parse(e.data).type ?? "unknown");
|
||||
} catch {
|
||||
log("info", "EVENT", e.data.slice(0, 100));
|
||||
}
|
||||
};
|
||||
|
||||
// Mic
|
||||
try {
|
||||
const stream = await navigator.mediaDevices.getUserMedia({ audio: true });
|
||||
streamRef.current = stream;
|
||||
stream.getTracks().forEach((t) => pc.addTrack(t, stream));
|
||||
log("success", "MIC", "Microphone access granted");
|
||||
setAudioStatus("Mic active — waiting for remote audio");
|
||||
micRef.current = true;
|
||||
setMicActive(true);
|
||||
} catch (e) {
|
||||
log("warn", "MIC", `Mic denied: ${e.message}`);
|
||||
const ctx = new AudioContext();
|
||||
const dest = ctx.createMediaStreamDestination();
|
||||
dest.stream.getTracks().forEach((t) => pc.addTrack(t, dest.stream));
|
||||
}
|
||||
|
||||
// Step 3: SDP offer
|
||||
log("step", "STEP 3", "Creating SDP offer");
|
||||
const offer = await pc.createOffer();
|
||||
await pc.setLocalDescription(offer);
|
||||
setSdpOffer(offer.sdp);
|
||||
setOfferActive(true);
|
||||
log("info", "SDP", `Offer created (${offer.sdp.split("\n").length} lines)`);
|
||||
|
||||
// Step 4: SDP exchange
|
||||
setFlowStep(2);
|
||||
log("step", "STEP 4", `POST ${url}/v1/realtime/calls`);
|
||||
try {
|
||||
const r = await fetch(`${url}/v1/realtime/calls`, {
|
||||
method: "POST",
|
||||
headers: { Authorization: `Bearer ${token}`, "Content-Type": "application/sdp" },
|
||||
body: offer.sdp,
|
||||
});
|
||||
log("info", "HTTP", `${r.status} ${r.statusText}`);
|
||||
if (!r.ok) {
|
||||
log("error", "ERR", await r.text());
|
||||
stopSession();
|
||||
return;
|
||||
}
|
||||
const ans = await r.text();
|
||||
log("success", "SDP", `Answer received (${ans.split("\n").length} lines)`);
|
||||
|
||||
// Step 5: remote description
|
||||
log("step", "STEP 5", "Setting remote description");
|
||||
await pc.setRemoteDescription({ type: "answer", sdp: ans });
|
||||
setSdpAnswer(ans);
|
||||
setAnswerActive(true);
|
||||
log("success", "CONN", "✓ Session established — Browser ↔ LiteLLM ↔ OpenAI");
|
||||
} catch (e) {
|
||||
log("error", "ERR", `calls failed: ${e.message}`);
|
||||
stopSession();
|
||||
}
|
||||
}
|
||||
|
||||
function stopSession() {
|
||||
if (pcRef.current) {
|
||||
pcRef.current.close();
|
||||
pcRef.current = null;
|
||||
}
|
||||
if (streamRef.current) {
|
||||
streamRef.current.getTracks().forEach((t) => t.stop());
|
||||
streamRef.current = null;
|
||||
}
|
||||
if (animRef.current) {
|
||||
cancelAnimationFrame(animRef.current);
|
||||
animRef.current = null;
|
||||
}
|
||||
tokenRef.current = null;
|
||||
micRef.current = false;
|
||||
setConnected(false);
|
||||
setStatus("idle");
|
||||
setFlowStep(0);
|
||||
setTokenPreview("—");
|
||||
setIceState("—");
|
||||
setConnState("—");
|
||||
setDcState("—");
|
||||
setMicActive(false);
|
||||
setOfferActive(false);
|
||||
setAnswerActive(false);
|
||||
setBars(Array(28).fill(2));
|
||||
setAudioStatus("Start a session first");
|
||||
log("warn", "SESSION", "Session stopped");
|
||||
}
|
||||
|
||||
function toggleMic() {
|
||||
if (!streamRef.current) {
|
||||
log("warn", "MIC", "No active session");
|
||||
return;
|
||||
}
|
||||
const next = !micRef.current;
|
||||
micRef.current = next;
|
||||
streamRef.current.getAudioTracks().forEach((t) => {
|
||||
t.enabled = next;
|
||||
});
|
||||
setMicActive(next);
|
||||
log("info", "MIC", next ? "Unmuted" : "Muted");
|
||||
}
|
||||
|
||||
const f = (n) => flowStep >= n;
|
||||
|
||||
return (
|
||||
<>
|
||||
<style>{STYLES}</style>
|
||||
<div className="wrt-wrap">
|
||||
{/* Toggle header */}
|
||||
<div className={`wrt-toggle${open ? "" : " closed"}`} onClick={() => setOpen((o) => !o)}>
|
||||
<div className="wrt-toggle-left">
|
||||
<div className="wrt-live-dot" />
|
||||
<div>
|
||||
<div className="wrt-toggle-title">INTERACTIVE TESTER</div>
|
||||
<div className="wrt-toggle-sub">Browser → LiteLLM → OpenAI · WebRTC</div>
|
||||
</div>
|
||||
</div>
|
||||
<span className={`wrt-chevron${open ? " open" : ""}`}>▼</span>
|
||||
</div>
|
||||
|
||||
{open && (
|
||||
<div className="wrt-body">
|
||||
{/* Sidebar */}
|
||||
<div className="wrt-sidebar">
|
||||
<div>
|
||||
<div className="wrt-label">Proxy Config</div>
|
||||
<div className="wrt-field">
|
||||
<label>Proxy URL</label>
|
||||
<input
|
||||
value={proxyUrl}
|
||||
onChange={(e) => setProxyUrl(e.target.value)}
|
||||
placeholder="http://localhost:4000"
|
||||
/>
|
||||
</div>
|
||||
<div className="wrt-field">
|
||||
<label>API Key</label>
|
||||
<input
|
||||
type="password"
|
||||
value={apiKey}
|
||||
onChange={(e) => setApiKey(e.target.value)}
|
||||
placeholder="sk-1234"
|
||||
/>
|
||||
</div>
|
||||
<div className="wrt-field">
|
||||
<label>Model</label>
|
||||
<input value={model} onChange={(e) => setModel(e.target.value)} />
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="wrt-divider" />
|
||||
|
||||
<div>
|
||||
<div className="wrt-label">Flow</div>
|
||||
<div className="wrt-flow">
|
||||
<div className={`wrt-flow-box${f(1) ? " active" : ""}`}>Browser</div>
|
||||
<div className={`wrt-flow-arrow${f(1) ? " active" : ""}`}>→</div>
|
||||
<div className={`wrt-flow-box${f(1) ? " active" : ""}`}>LiteLLM</div>
|
||||
<div className={`wrt-flow-arrow${f(2) ? " active" : ""}`}>→</div>
|
||||
<div className={`wrt-flow-box${f(2) ? " active" : ""}`}>OpenAI</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="wrt-divider" />
|
||||
|
||||
<div>
|
||||
<div className="wrt-label">Controls</div>
|
||||
<button className="wrt-btn wrt-btn-primary" onClick={startSession} disabled={connected}>
|
||||
▶ Start Session
|
||||
</button>
|
||||
<button className="wrt-btn wrt-btn-danger" onClick={stopSession} disabled={!connected}>
|
||||
■ Stop
|
||||
</button>
|
||||
<button className="wrt-btn wrt-btn-ghost" onClick={clearLogs} style={{ marginTop: 5 }}>
|
||||
✕ Clear Logs
|
||||
</button>
|
||||
</div>
|
||||
|
||||
<div className="wrt-divider" />
|
||||
|
||||
<div>
|
||||
<div className="wrt-label">Session Info</div>
|
||||
<div className="wrt-meta">
|
||||
{[
|
||||
["token", tokenPreview],
|
||||
["ice", iceState],
|
||||
["conn", connState],
|
||||
["data ch.", dcState],
|
||||
].map(([k, v]) => (
|
||||
<div className="wrt-meta-row" key={k}>
|
||||
<span>{k}</span>
|
||||
<span>{v}</span>
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* Right panel */}
|
||||
<div className="wrt-main">
|
||||
<div className="wrt-header">
|
||||
<span className="wrt-header-title">WEBRTC REALTIME TESTER</span>
|
||||
<div className="wrt-status-pill">
|
||||
<div className={`wrt-status-dot${status !== "idle" ? ` ${status}` : ""}`} />
|
||||
<span style={{ fontSize: 10, color: "#4a5568" }}>{status}</span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="wrt-tabs">
|
||||
{["logs", "sdp", "audio"].map((t) => (
|
||||
<div key={t} className={`wrt-tab${activeTab === t ? " active" : ""}`} onClick={() => setActiveTab(t)}>
|
||||
{t.toUpperCase()}
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
|
||||
{/* Logs */}
|
||||
<div className={`wrt-tab-content${activeTab === "logs" ? " active" : ""}`}>
|
||||
<div className="wrt-log" ref={logRef}>
|
||||
{entries.length === 0 ? (
|
||||
<div className="wrt-empty">
|
||||
<div style={{ fontSize: 22, opacity: 0.3 }}>📡</div>
|
||||
<div>Hit "Start Session" to begin</div>
|
||||
</div>
|
||||
) : (
|
||||
entries.map((e) => (
|
||||
<div key={e.id} className={`wrt-entry ${e.level}`}>
|
||||
<span className="we-time">{e.time}</span>
|
||||
<span className="we-tag">[{e.tag}]</span>
|
||||
<span className="we-msg">{e.msg}</span>
|
||||
</div>
|
||||
))
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* SDP */}
|
||||
<div className={`wrt-tab-content${activeTab === "sdp" ? " active" : ""}`}>
|
||||
<div className="wrt-sdp-pane">
|
||||
<div className="wrt-sdp-box">
|
||||
<div className="wrt-sdp-hdr">
|
||||
<div className={`wrt-sdp-dot${offerActive ? " active" : ""}`} />
|
||||
SDP OFFER
|
||||
</div>
|
||||
<textarea readOnly value={sdpOffer} placeholder="SDP offer appears here..." />
|
||||
</div>
|
||||
<div className="wrt-sdp-box">
|
||||
<div className="wrt-sdp-hdr">
|
||||
<div className={`wrt-sdp-dot${answerActive ? " active" : ""}`} />
|
||||
SDP ANSWER
|
||||
</div>
|
||||
<textarea readOnly value={sdpAnswer} placeholder="SDP answer appears here..." />
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* Audio */}
|
||||
<div className={`wrt-tab-content${activeTab === "audio" ? " active" : ""}`}>
|
||||
<div className="wrt-audio-pane">
|
||||
<div className="wrt-viz">
|
||||
{bars.map((h, i) => (
|
||||
<div
|
||||
key={i}
|
||||
className="wrt-bar"
|
||||
style={{ height: h + "px", background: `hsl(${150 - (h / 42) * 30},100%,55%)` }}
|
||||
/>
|
||||
))}
|
||||
</div>
|
||||
<button className={`wrt-mic-btn${micActive ? " active" : ""}`} onClick={toggleMic}>
|
||||
🎙️
|
||||
</button>
|
||||
<div className="wrt-audio-status">{audioStatus}</div>
|
||||
<audio ref={remoteAudioRef} autoPlay style={{ display: "none" }} />
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</>
|
||||
);
|
||||
}
|
||||
|
|
@ -14,12 +14,6 @@ vi.mock("./agents/add_agent_form", () => ({
|
|||
default: () => <div data-testid="add-agent-form" />,
|
||||
}));
|
||||
|
||||
vi.mock("./agents/agent_card_grid", () => ({
|
||||
default: ({ isAdmin }: { isAdmin: boolean }) => <div data-testid="agent-card-grid" data-is-admin={String(isAdmin)} />,
|
||||
}));
|
||||
|
||||
// Note: agents.tsx no longer uses AgentCardGrid — it renders a Table directly.
|
||||
|
||||
vi.mock("./agents/agent_info", () => ({
|
||||
default: () => <div data-testid="agent-info" />,
|
||||
}));
|
||||
|
|
|
|||
|
|
@ -1,95 +0,0 @@
|
|||
import React from "react";
|
||||
import { describe, it, expect, vi } from "vitest";
|
||||
import { screen } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { renderWithProviders } from "../../../tests/test-utils";
|
||||
import AgentCard from "./agent_card";
|
||||
import type { Agent } from "./types";
|
||||
|
||||
const baseAgent: Agent = {
|
||||
agent_id: "agent-123",
|
||||
agent_name: "Test Agent",
|
||||
litellm_params: { model: "gpt-4" },
|
||||
agent_card_params: {
|
||||
description: "A test agent for unit testing",
|
||||
url: "https://agent.example.com",
|
||||
},
|
||||
};
|
||||
|
||||
const defaultProps = {
|
||||
agent: baseAgent,
|
||||
onAgentClick: vi.fn(),
|
||||
accessToken: "token-123",
|
||||
isAdmin: false,
|
||||
onAgentUpdated: vi.fn(),
|
||||
};
|
||||
|
||||
describe("AgentCard", () => {
|
||||
it("should render the agent name and description", () => {
|
||||
renderWithProviders(<AgentCard {...defaultProps} />);
|
||||
|
||||
expect(screen.getByText("Test Agent")).toBeInTheDocument();
|
||||
expect(screen.getByText("A test agent for unit testing")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show 'No description' when agent has no description", () => {
|
||||
const agent = { ...baseAgent, agent_card_params: {} };
|
||||
renderWithProviders(<AgentCard {...defaultProps} agent={agent} />);
|
||||
|
||||
expect(screen.getByText("No description")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show the agent URL when provided", () => {
|
||||
renderWithProviders(<AgentCard {...defaultProps} />);
|
||||
|
||||
expect(screen.getByText("https://agent.example.com")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show 'Needs Setup' badge when agent has no key", () => {
|
||||
renderWithProviders(<AgentCard {...defaultProps} />);
|
||||
|
||||
expect(screen.getByText("Needs Setup")).toBeInTheDocument();
|
||||
expect(screen.getByText("No key assigned")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show 'Active' badge and key info when agent has a key", () => {
|
||||
const keyInfo = { has_key: true, key_alias: "my-key" };
|
||||
renderWithProviders(<AgentCard {...defaultProps} keyInfo={keyInfo} />);
|
||||
|
||||
expect(screen.getByText("Active")).toBeInTheDocument();
|
||||
expect(screen.getByText("my-key")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should call onAgentClick when card is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onAgentClick = vi.fn();
|
||||
renderWithProviders(<AgentCard {...defaultProps} onAgentClick={onAgentClick} />);
|
||||
|
||||
await user.click(screen.getByText("Test Agent"));
|
||||
|
||||
expect(onAgentClick).toHaveBeenCalledWith("agent-123");
|
||||
});
|
||||
|
||||
it("should show delete button only for admins", () => {
|
||||
const onDeleteClick = vi.fn();
|
||||
const { unmount } = renderWithProviders(
|
||||
<AgentCard {...defaultProps} isAdmin={false} onDeleteClick={onDeleteClick} />,
|
||||
);
|
||||
expect(screen.queryByRole("button", { name: /delete/i })).not.toBeInTheDocument();
|
||||
|
||||
unmount();
|
||||
|
||||
renderWithProviders(<AgentCard {...defaultProps} isAdmin={true} onDeleteClick={onDeleteClick} />);
|
||||
expect(screen.getByRole("button", { name: /delete/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should call onDeleteClick with agent id and name when delete is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
const onDeleteClick = vi.fn();
|
||||
renderWithProviders(<AgentCard {...defaultProps} isAdmin={true} onDeleteClick={onDeleteClick} />);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /delete/i }));
|
||||
|
||||
expect(onDeleteClick).toHaveBeenCalledWith("agent-123", "Test Agent");
|
||||
});
|
||||
});
|
||||
|
|
@ -1,88 +0,0 @@
|
|||
import React from "react";
|
||||
import { Card, Badge, Tooltip, Button } from "antd";
|
||||
import { CopyOutlined, KeyOutlined, WarningOutlined, DeleteOutlined } from "@ant-design/icons";
|
||||
import { Agent, AgentKeyInfo } from "./types";
|
||||
|
||||
interface AgentCardProps {
|
||||
agent: Agent;
|
||||
keyInfo?: AgentKeyInfo;
|
||||
onAgentClick: (agentId: string) => void;
|
||||
onDeleteClick?: (agentId: string, agentName: string) => void;
|
||||
accessToken: string | null;
|
||||
isAdmin: boolean;
|
||||
onAgentUpdated: () => void;
|
||||
}
|
||||
|
||||
const AgentCard: React.FC<AgentCardProps> = ({ agent, keyInfo, onAgentClick, onDeleteClick, isAdmin }) => {
|
||||
const description = agent.agent_card_params?.description || "No description";
|
||||
const url = agent.agent_card_params?.url;
|
||||
const hasKey = keyInfo?.has_key ?? false;
|
||||
const statusBadge = hasKey ? <Badge status="success" text="Active" /> : <Badge status="warning" text="Needs Setup" />;
|
||||
|
||||
const copyToClipboard = (e: React.MouseEvent, text: string) => {
|
||||
e.stopPropagation();
|
||||
navigator.clipboard.writeText(text);
|
||||
};
|
||||
|
||||
return (
|
||||
<Card
|
||||
hoverable
|
||||
className="h-full flex flex-col"
|
||||
styles={{
|
||||
body: { flex: 1, display: "flex", flexDirection: "column" },
|
||||
}}
|
||||
onClick={() => onAgentClick(agent.agent_id)}
|
||||
>
|
||||
<div className="flex items-start justify-between gap-2 mb-2">
|
||||
<div className="flex-1 min-w-0">
|
||||
<div className="flex items-center gap-2 flex-wrap">
|
||||
<span className="font-medium text-gray-900 truncate">{agent.agent_name}</span>
|
||||
<Tooltip title="Copy Agent ID">
|
||||
<CopyOutlined
|
||||
onClick={(e) => copyToClipboard(e, agent.agent_id)}
|
||||
className="cursor-pointer text-gray-400 hover:text-blue-500 text-xs shrink-0"
|
||||
/>
|
||||
</Tooltip>
|
||||
</div>
|
||||
<div className="mt-1">{statusBadge}</div>
|
||||
</div>
|
||||
{isAdmin && onDeleteClick && (
|
||||
<Tooltip title="Delete agent">
|
||||
<Button
|
||||
type="text"
|
||||
size="small"
|
||||
danger
|
||||
icon={<DeleteOutlined />}
|
||||
onClick={(e) => {
|
||||
e.stopPropagation();
|
||||
onDeleteClick(agent.agent_id, agent.agent_name);
|
||||
}}
|
||||
className="shrink-0 -mr-1"
|
||||
/>
|
||||
</Tooltip>
|
||||
)}
|
||||
</div>
|
||||
<p className="text-sm text-gray-600 line-clamp-2 flex-1 mb-3">{description}</p>
|
||||
{url && (
|
||||
<p className="text-xs text-gray-500 truncate mb-2" title={url}>
|
||||
{url}
|
||||
</p>
|
||||
)}
|
||||
<div className="mt-auto pt-3 border-t border-gray-100 text-xs">
|
||||
{hasKey ? (
|
||||
<div className="flex items-center gap-1.5 text-gray-600">
|
||||
<KeyOutlined />
|
||||
<span>{keyInfo?.key_alias || keyInfo?.token_prefix || "Key assigned"}</span>
|
||||
</div>
|
||||
) : (
|
||||
<div className="flex items-center gap-1.5 text-amber-600">
|
||||
<WarningOutlined />
|
||||
<span>No key assigned</span>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</Card>
|
||||
);
|
||||
};
|
||||
|
||||
export default AgentCard;
|
||||
|
|
@ -1,80 +0,0 @@
|
|||
import { renderWithProviders, screen } from "../../../tests/test-utils";
|
||||
import { vi } from "vitest";
|
||||
import AgentCardGrid from "./agent_card_grid";
|
||||
import type { Agent, AgentKeyInfo } from "./types";
|
||||
|
||||
vi.mock("./agent_card", () => ({
|
||||
default: ({ agent, onAgentClick }: any) => (
|
||||
<div data-testid={`agent-card-${agent.agent_id}`} onClick={() => onAgentClick(agent.agent_id)}>
|
||||
{agent.agent_name}
|
||||
</div>
|
||||
),
|
||||
}));
|
||||
|
||||
const mockAgents: Agent[] = [
|
||||
{
|
||||
agent_id: "agent-1",
|
||||
agent_name: "Test Agent 1",
|
||||
litellm_params: { model: "gpt-4" },
|
||||
agent_card_params: { description: "First agent" },
|
||||
},
|
||||
{
|
||||
agent_id: "agent-2",
|
||||
agent_name: "Test Agent 2",
|
||||
litellm_params: { model: "claude-3" },
|
||||
agent_card_params: { description: "Second agent" },
|
||||
},
|
||||
];
|
||||
|
||||
const mockKeyInfoMap: Record<string, AgentKeyInfo> = {
|
||||
"agent-1": { has_key: true, key_alias: "key-1" },
|
||||
"agent-2": { has_key: false },
|
||||
};
|
||||
|
||||
const defaultProps = {
|
||||
agentsList: mockAgents,
|
||||
keyInfoMap: mockKeyInfoMap,
|
||||
isLoading: false,
|
||||
onDeleteClick: vi.fn(),
|
||||
accessToken: "test-token",
|
||||
onAgentUpdated: vi.fn(),
|
||||
isAdmin: true,
|
||||
onAgentClick: vi.fn(),
|
||||
};
|
||||
|
||||
describe("AgentCardGrid", () => {
|
||||
it("should render", () => {
|
||||
renderWithProviders(<AgentCardGrid {...defaultProps} />);
|
||||
expect(screen.getByText("Test Agent 1")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render all agent cards", () => {
|
||||
renderWithProviders(<AgentCardGrid {...defaultProps} />);
|
||||
expect(screen.getByText("Test Agent 1")).toBeInTheDocument();
|
||||
expect(screen.getByText("Test Agent 2")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show loading skeletons when isLoading is true", () => {
|
||||
renderWithProviders(<AgentCardGrid {...defaultProps} isLoading={true} />);
|
||||
expect(screen.queryByText("Test Agent 1")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show admin empty state message when no agents and isAdmin", () => {
|
||||
renderWithProviders(<AgentCardGrid {...defaultProps} agentsList={[]} isAdmin={true} />);
|
||||
expect(screen.getByText("No agents found. Create one to get started.")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show non-admin empty state message when no agents and not admin", () => {
|
||||
renderWithProviders(<AgentCardGrid {...defaultProps} agentsList={[]} isAdmin={false} />);
|
||||
expect(screen.getByText("No agents found. Contact an admin to create agents.")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should call onAgentClick when a card is clicked", async () => {
|
||||
const onAgentClick = vi.fn();
|
||||
renderWithProviders(<AgentCardGrid {...defaultProps} onAgentClick={onAgentClick} />);
|
||||
const { default: userEvent } = await import("@testing-library/user-event");
|
||||
const user = userEvent.setup();
|
||||
await user.click(screen.getByTestId("agent-card-agent-1"));
|
||||
expect(onAgentClick).toHaveBeenCalledWith("agent-1");
|
||||
});
|
||||
});
|
||||
|
|
@ -1,67 +0,0 @@
|
|||
import React from "react";
|
||||
import { Skeleton } from "antd";
|
||||
import AgentCard from "./agent_card";
|
||||
import { Agent, AgentKeyInfo } from "./types";
|
||||
|
||||
interface AgentCardGridProps {
|
||||
agentsList: Agent[];
|
||||
keyInfoMap: Record<string, AgentKeyInfo>;
|
||||
isLoading: boolean;
|
||||
onDeleteClick: (agentId: string, agentName: string) => void;
|
||||
accessToken: string | null;
|
||||
onAgentUpdated: () => void;
|
||||
isAdmin: boolean;
|
||||
onAgentClick: (agentId: string) => void;
|
||||
}
|
||||
|
||||
const AgentCardGrid: React.FC<AgentCardGridProps> = ({
|
||||
agentsList,
|
||||
keyInfoMap,
|
||||
isLoading,
|
||||
onDeleteClick,
|
||||
accessToken,
|
||||
onAgentUpdated,
|
||||
isAdmin,
|
||||
onAgentClick,
|
||||
}) => {
|
||||
if (isLoading) {
|
||||
return (
|
||||
<div className="grid grid-cols-1 md:grid-cols-2 lg:grid-cols-3 gap-6">
|
||||
{[1, 2, 3].map((i) => (
|
||||
<Skeleton key={i} active paragraph={{ rows: 3 }} />
|
||||
))}
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
if (!agentsList || agentsList.length === 0) {
|
||||
return (
|
||||
<div className="rounded-lg border border-gray-200 bg-gray-50/50 py-12 text-center">
|
||||
<p className="text-gray-500">
|
||||
{isAdmin
|
||||
? "No agents found. Create one to get started."
|
||||
: "No agents found. Contact an admin to create agents."}
|
||||
</p>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
return (
|
||||
<div className="grid grid-cols-1 md:grid-cols-2 lg:grid-cols-3 gap-6">
|
||||
{agentsList.map((agent) => (
|
||||
<AgentCard
|
||||
key={agent.agent_id}
|
||||
agent={agent}
|
||||
keyInfo={keyInfoMap[agent.agent_id]}
|
||||
onAgentClick={onAgentClick}
|
||||
onDeleteClick={isAdmin ? onDeleteClick : undefined}
|
||||
accessToken={accessToken}
|
||||
isAdmin={isAdmin}
|
||||
onAgentUpdated={onAgentUpdated}
|
||||
/>
|
||||
))}
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
export default AgentCardGrid;
|
||||
|
|
@ -210,6 +210,12 @@ export const GUARDRAIL_PRESETS: Record<string, GuardrailPreset> = {
|
|||
mode: "pre_call",
|
||||
defaultOn: false,
|
||||
},
|
||||
cisco_ai_defense: {
|
||||
provider: "CiscoAiDefense",
|
||||
guardrailNameSuggestion: "Cisco AI Defense",
|
||||
mode: "pre_call",
|
||||
defaultOn: false,
|
||||
},
|
||||
noma: {
|
||||
provider: "Noma",
|
||||
guardrailNameSuggestion: "Noma Security",
|
||||
|
|
|
|||
|
|
@ -307,6 +307,16 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [
|
|||
logo: `${ASSET_PREFIX}palo_alto_networks.jpeg`,
|
||||
tags: ["Enterprise", "Security"],
|
||||
},
|
||||
{
|
||||
id: "cisco_ai_defense",
|
||||
name: "Cisco AI Defense",
|
||||
description:
|
||||
"Cisco AI Defense Inspection API for runtime protection: prompt injection, PII/PCI/PHI, harassment, hate speech, profanity, violence, and code detection.",
|
||||
category: "partner",
|
||||
logo: `${ASSET_PREFIX}cisco.png`,
|
||||
tags: ["Enterprise", "Security", "Prompt Injection", "PII"],
|
||||
providerKey: "CiscoAiDefense",
|
||||
},
|
||||
{
|
||||
id: "noma",
|
||||
name: "Noma Security",
|
||||
|
|
|
|||
|
|
@ -123,6 +123,7 @@ export const guardrailLogoMap: Record<string, string> = {
|
|||
"Azure Content Safety Text Moderation": `${asset_logos_folder}microsoft_azure.svg`,
|
||||
"Aporia AI": `${asset_logos_folder}aporia.png`,
|
||||
"PANW Prisma AIRS": `${asset_logos_folder}palo_alto_networks.jpeg`,
|
||||
"Cisco AI Defense": `${asset_logos_folder}cisco.png`,
|
||||
"Noma Security": `${asset_logos_folder}noma_security.png`,
|
||||
"Javelin Guardrails": `${asset_logos_folder}javelin.png`,
|
||||
"Pillar Guardrail": `${asset_logos_folder}pillar.jpeg`,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
|
||||
import { render } from "@testing-library/react";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
import React from "react";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
|
||||
vi.mock("./vector_store_management/VectorStoreSelector", () => ({
|
||||
__esModule: true,
|
||||
|
|
@ -10,22 +11,27 @@ vi.mock("./mcp_server_management/MCPServerSelector", () => ({
|
|||
__esModule: true,
|
||||
default: () => null,
|
||||
}));
|
||||
vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
|
||||
default: () => ({
|
||||
accessToken: null,
|
||||
userId: null,
|
||||
userRole: null,
|
||||
}),
|
||||
}));
|
||||
|
||||
import OrganizationsTable from "./organizations";
|
||||
|
||||
const renderWithQueryClient = (ui: React.ReactElement) => {
|
||||
const queryClient = new QueryClient({
|
||||
defaultOptions: { queries: { retry: false } },
|
||||
});
|
||||
return render(<QueryClientProvider client={queryClient}>{ui}</QueryClientProvider>);
|
||||
};
|
||||
|
||||
describe("OrganizationsTable", () => {
|
||||
it("should render the OrganizationsTable component", () => {
|
||||
const setOrganizations = vi.fn();
|
||||
|
||||
const { getByText } = render(
|
||||
<OrganizationsTable
|
||||
organizations={[]}
|
||||
userRole="Admin"
|
||||
userModels={[]}
|
||||
accessToken={null}
|
||||
setOrganizations={setOrganizations}
|
||||
premiumUser={true}
|
||||
/>,
|
||||
const { getByText } = renderWithQueryClient(
|
||||
<OrganizationsTable userRole="Admin" accessToken={null} premiumUser={true} />,
|
||||
);
|
||||
|
||||
expect(getByText("+ Create New Organization")).toBeInTheDocument();
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
import { organizationKeys, useOrganizations } from "@/app/(dashboard)/hooks/organizations/useOrganizations";
|
||||
import { useUserModels } from "@/app/(dashboard)/hooks/models/useModels";
|
||||
import OrganizationFilters, { FilterState } from "@/app/(dashboard)/organizations/OrganizationFilters";
|
||||
import { InfoCircleOutlined } from "@ant-design/icons";
|
||||
import { ChevronDownIcon, ChevronRightIcon, RefreshIcon } from "@heroicons/react/outline";
|
||||
|
|
@ -23,6 +25,7 @@ import {
|
|||
TextInput,
|
||||
} from "@tremor/react";
|
||||
import { Form, Input, Modal, Select as Select2, Tooltip } from "antd";
|
||||
import { useQueryClient } from "@tanstack/react-query";
|
||||
import React, { useState } from "react";
|
||||
import { formatNumberWithCommas } from "../utils/dataUtils";
|
||||
import DeleteResourceModal from "./common_components/DeleteResourceModal";
|
||||
|
|
@ -37,15 +40,10 @@ import NumericalInput from "./shared/numerical_input";
|
|||
import VectorStoreSelector from "./vector_store_management/VectorStoreSelector";
|
||||
|
||||
interface OrganizationsTableProps {
|
||||
organizations: Organization[];
|
||||
userRole: string;
|
||||
userModels: string[];
|
||||
accessToken: string | null;
|
||||
lastRefreshed?: string;
|
||||
handleRefreshClick?: () => void;
|
||||
currentOrg?: any;
|
||||
guardrailsList?: string[];
|
||||
setOrganizations: (organizations: Organization[]) => void;
|
||||
premiumUser: boolean;
|
||||
}
|
||||
|
||||
|
|
@ -60,15 +58,10 @@ export const fetchOrganizations = async (
|
|||
};
|
||||
|
||||
const OrganizationsTable: React.FC<OrganizationsTableProps> = ({
|
||||
organizations,
|
||||
userRole,
|
||||
userModels,
|
||||
accessToken,
|
||||
lastRefreshed,
|
||||
handleRefreshClick,
|
||||
currentOrg,
|
||||
guardrailsList = [],
|
||||
setOrganizations,
|
||||
premiumUser,
|
||||
}) => {
|
||||
const [selectedOrgId, setSelectedOrgId] = useState<string | null>(null);
|
||||
|
|
@ -87,21 +80,14 @@ const OrganizationsTable: React.FC<OrganizationsTableProps> = ({
|
|||
sort_order: "desc",
|
||||
});
|
||||
|
||||
const queryClient = useQueryClient();
|
||||
const { data: organizations = [] } = useOrganizations({ org_id: filters.org_id, org_alias: filters.org_alias });
|
||||
const { data: userModels = [] } = useUserModels();
|
||||
|
||||
const refetchOrganizations = () => queryClient.invalidateQueries({ queryKey: organizationKeys.lists() });
|
||||
|
||||
const handleFilterChange = (key: keyof FilterState, value: string) => {
|
||||
const newFilters = { ...filters, [key]: value };
|
||||
setFilters(newFilters);
|
||||
// Call organizationListCall with the new filters
|
||||
if (accessToken) {
|
||||
organizationListCall(accessToken, newFilters.org_id || null, newFilters.org_alias || null)
|
||||
.then((response) => {
|
||||
if (response) {
|
||||
setOrganizations(response);
|
||||
}
|
||||
})
|
||||
.catch((error) => {
|
||||
console.error("Error fetching organizations:", error);
|
||||
});
|
||||
}
|
||||
setFilters((previousFilters) => ({ ...previousFilters, [key]: value }));
|
||||
};
|
||||
|
||||
const handleFilterReset = () => {
|
||||
|
|
@ -111,18 +97,6 @@ const OrganizationsTable: React.FC<OrganizationsTableProps> = ({
|
|||
sort_by: "created_at",
|
||||
sort_order: "desc",
|
||||
});
|
||||
// Reset organizations list
|
||||
if (accessToken) {
|
||||
organizationListCall(accessToken, null, null)
|
||||
.then((response) => {
|
||||
if (response) {
|
||||
setOrganizations(response);
|
||||
}
|
||||
})
|
||||
.catch((error) => {
|
||||
console.error("Error fetching organizations:", error);
|
||||
});
|
||||
}
|
||||
};
|
||||
|
||||
const handleDelete = (orgId: string | null) => {
|
||||
|
|
@ -142,8 +116,7 @@ const OrganizationsTable: React.FC<OrganizationsTableProps> = ({
|
|||
|
||||
setIsDeleteModalOpen(false);
|
||||
setOrgToDelete(null);
|
||||
// Refresh organizations list
|
||||
await fetchOrganizations(accessToken, setOrganizations, filters.org_id || null, filters.org_alias || null);
|
||||
await refetchOrganizations();
|
||||
} catch (error) {
|
||||
console.error("Error deleting organization:", error);
|
||||
} finally {
|
||||
|
|
@ -189,8 +162,7 @@ const OrganizationsTable: React.FC<OrganizationsTableProps> = ({
|
|||
NotificationsManager.success("Organization created successfully");
|
||||
setIsOrgModalVisible(false);
|
||||
form.resetFields();
|
||||
// Refresh organizations list
|
||||
fetchOrganizations(accessToken, setOrganizations, filters.org_id || null, filters.org_alias || null);
|
||||
await refetchOrganizations();
|
||||
} catch (error) {
|
||||
console.error("Error creating organization:", error);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,85 +0,0 @@
|
|||
import { render, screen } from "@testing-library/react";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { ErrorViewer } from "./ErrorViewer";
|
||||
|
||||
const basicError = {
|
||||
error_class: "NotFoundError",
|
||||
error_message: "Model gpt-5 not found",
|
||||
};
|
||||
|
||||
const errorWithTraceback = {
|
||||
error_class: "AuthenticationError",
|
||||
error_message: "Invalid API key",
|
||||
traceback: `Traceback (most recent call last):
|
||||
File "/app/main.py", line 42, in handle_request
|
||||
result = await client.chat(model="gpt-4")
|
||||
File "/app/llms/openai.py", line 100, in chat
|
||||
response = self._make_request(payload)
|
||||
File "/app/llms/base.py", line 55, in _make_request
|
||||
raise AuthenticationError("Invalid API key")`,
|
||||
};
|
||||
|
||||
describe("ErrorViewer", () => {
|
||||
it("should render error type and message", () => {
|
||||
render(<ErrorViewer errorInfo={basicError} />);
|
||||
expect(screen.getByText("NotFoundError")).toBeInTheDocument();
|
||||
expect(screen.getByText("Model gpt-5 not found")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show 'Unknown Error' when error_class is missing", () => {
|
||||
render(<ErrorViewer errorInfo={{ error_message: "something broke" }} />);
|
||||
expect(screen.getByText("Unknown Error")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show 'Unknown error occurred' when error_message is missing", () => {
|
||||
render(<ErrorViewer errorInfo={{ error_class: "RuntimeError" }} />);
|
||||
expect(screen.getByText("Unknown error occurred")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render traceback frames when traceback is present", () => {
|
||||
render(<ErrorViewer errorInfo={errorWithTraceback} />);
|
||||
expect(screen.getByText("Traceback")).toBeInTheDocument();
|
||||
expect(screen.getByText("main.py")).toBeInTheDocument();
|
||||
expect(screen.getByText("openai.py")).toBeInTheDocument();
|
||||
expect(screen.getByText("base.py")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should expand a frame when clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<ErrorViewer errorInfo={errorWithTraceback} />);
|
||||
|
||||
await user.click(screen.getByText("main.py"));
|
||||
|
||||
expect(screen.getByText('result = await client.chat(model="gpt-4")')).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should expand all frames when 'Expand All' is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<ErrorViewer errorInfo={errorWithTraceback} />);
|
||||
|
||||
await user.click(screen.getByText("Expand All"));
|
||||
|
||||
expect(screen.getByText("Collapse All")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should copy traceback to clipboard when copy button is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
const mockWriteText = vi.fn().mockResolvedValue(undefined);
|
||||
Object.defineProperty(navigator, "clipboard", {
|
||||
value: { writeText: mockWriteText },
|
||||
writable: true,
|
||||
configurable: true,
|
||||
});
|
||||
|
||||
render(<ErrorViewer errorInfo={errorWithTraceback} />);
|
||||
|
||||
await user.click(screen.getByTitle("Copy traceback"));
|
||||
expect(mockWriteText).toHaveBeenCalledWith(errorWithTraceback.traceback);
|
||||
});
|
||||
|
||||
it("should not render traceback section when traceback is absent", () => {
|
||||
render(<ErrorViewer errorInfo={basicError} />);
|
||||
expect(screen.queryByText("Traceback")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
@ -1,180 +0,0 @@
|
|||
import React from "react";
|
||||
|
||||
interface ErrorViewerProps {
|
||||
errorInfo: {
|
||||
error_class?: string;
|
||||
error_message?: string;
|
||||
traceback?: string;
|
||||
llm_provider?: string;
|
||||
error_code?: string | number;
|
||||
};
|
||||
}
|
||||
|
||||
export const ErrorViewer: React.FC<ErrorViewerProps> = ({ errorInfo }) => {
|
||||
const [expandedFrames, setExpandedFrames] = React.useState<{ [key: number]: boolean }>({});
|
||||
const [allExpanded, setAllExpanded] = React.useState(false);
|
||||
|
||||
// Toggle individual frame
|
||||
const toggleFrame = (index: number) => {
|
||||
setExpandedFrames((prev) => ({
|
||||
...prev,
|
||||
[index]: !prev[index],
|
||||
}));
|
||||
};
|
||||
|
||||
// Toggle all frames
|
||||
const toggleAllFrames = () => {
|
||||
const newState = !allExpanded;
|
||||
setAllExpanded(newState);
|
||||
|
||||
if (tracebackFrames.length > 0) {
|
||||
const newExpandedState: { [key: number]: boolean } = {};
|
||||
tracebackFrames.forEach((_, idx) => {
|
||||
newExpandedState[idx] = newState;
|
||||
});
|
||||
setExpandedFrames(newExpandedState);
|
||||
}
|
||||
};
|
||||
|
||||
// Parse traceback into frames
|
||||
const parseTraceback = (traceback: string) => {
|
||||
if (!traceback) return [];
|
||||
|
||||
// Extract file paths, line numbers and code from traceback
|
||||
const fileLineRegex = /File "([^"]+)", line (\d+)/g;
|
||||
const matches = Array.from(traceback.matchAll(fileLineRegex));
|
||||
|
||||
// Create simplified frames
|
||||
return matches.map((match) => {
|
||||
const filePath = match[1];
|
||||
const lineNumber = match[2];
|
||||
const fileName = filePath.split("/").pop() || filePath;
|
||||
|
||||
// Extract the context around this frame
|
||||
const matchIndex = match.index || 0;
|
||||
const nextMatchIndex = traceback.indexOf('File "', matchIndex + 1);
|
||||
const frameContent =
|
||||
nextMatchIndex > -1
|
||||
? traceback.substring(matchIndex, nextMatchIndex).trim()
|
||||
: traceback.substring(matchIndex).trim();
|
||||
|
||||
// Try to extract the code line
|
||||
const lines = frameContent.split("\n");
|
||||
let code = "";
|
||||
if (lines.length > 1) {
|
||||
code = lines[lines.length - 1].trim();
|
||||
}
|
||||
|
||||
return {
|
||||
filePath,
|
||||
fileName,
|
||||
lineNumber,
|
||||
code,
|
||||
inFunction: frameContent.includes(" in ") ? frameContent.split(" in ")[1].split("\n")[0] : "",
|
||||
};
|
||||
});
|
||||
};
|
||||
|
||||
const tracebackFrames = errorInfo.traceback ? parseTraceback(errorInfo.traceback) : [];
|
||||
|
||||
return (
|
||||
<div className="bg-white rounded-lg shadow">
|
||||
<div className="p-4 border-b">
|
||||
<h3 className="text-lg font-medium flex items-center text-red-600">
|
||||
<svg className="w-5 h-5 mr-2" fill="none" stroke="currentColor" viewBox="0 0 24 24">
|
||||
<path
|
||||
strokeLinecap="round"
|
||||
strokeLinejoin="round"
|
||||
strokeWidth={2}
|
||||
d="M12 9v2m0 4h.01m-6.938 4h13.856c1.54 0 2.502-1.667 1.732-3L13.732 4c-.77-1.333-2.694-1.333-3.464 0L3.34 16c-.77 1.333.192 3 1.732 3z"
|
||||
/>
|
||||
</svg>
|
||||
Error Details
|
||||
</h3>
|
||||
</div>
|
||||
|
||||
<div className="p-4">
|
||||
<div className="bg-red-50 rounded-md p-4 mb-4">
|
||||
<div className="flex">
|
||||
<span className="text-red-800 font-medium w-20">Type:</span>
|
||||
<span className="text-red-700">{errorInfo.error_class || "Unknown Error"}</span>
|
||||
</div>
|
||||
<div className="flex mt-2">
|
||||
<span className="text-red-800 font-medium w-20 flex-shrink-0">Message:</span>
|
||||
<span className="text-red-700 break-words whitespace-pre-wrap">
|
||||
{errorInfo.error_message || "Unknown error occurred"}
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{errorInfo.traceback && (
|
||||
<div className="mt-4">
|
||||
<div className="flex justify-between items-center mb-2">
|
||||
<h4 className="font-medium">Traceback</h4>
|
||||
<div className="flex items-center space-x-4">
|
||||
<button
|
||||
onClick={toggleAllFrames}
|
||||
className="text-gray-500 hover:text-gray-700 flex items-center text-sm"
|
||||
>
|
||||
{allExpanded ? "Collapse All" : "Expand All"}
|
||||
</button>
|
||||
<button
|
||||
onClick={() => navigator.clipboard.writeText(errorInfo.traceback || "")}
|
||||
className="text-gray-500 hover:text-gray-700 flex items-center"
|
||||
title="Copy traceback"
|
||||
>
|
||||
<svg
|
||||
xmlns="http://www.w3.org/2000/svg"
|
||||
width="16"
|
||||
height="16"
|
||||
viewBox="0 0 24 24"
|
||||
fill="none"
|
||||
stroke="currentColor"
|
||||
strokeWidth="2"
|
||||
strokeLinecap="round"
|
||||
strokeLinejoin="round"
|
||||
>
|
||||
<rect x="9" y="9" width="13" height="13" rx="2" ry="2"></rect>
|
||||
<path d="M5 15H4a2 2 0 0 1-2-2V4a2 2 0 0 1 2-2h9a2 2 0 0 1 2 2v1"></path>
|
||||
</svg>
|
||||
<span className="ml-1">Copy</span>
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div className="bg-white rounded-md border border-gray-200 overflow-hidden shadow-sm">
|
||||
{tracebackFrames.map((frame, index) => (
|
||||
<div key={index} className="border-b border-gray-200 last:border-b-0">
|
||||
<div
|
||||
className="px-4 py-2 flex items-center justify-between cursor-pointer hover:bg-gray-50"
|
||||
onClick={() => toggleFrame(index)}
|
||||
>
|
||||
<div className="flex items-center">
|
||||
<span className="text-gray-400 mr-2 w-12 text-right">{frame.lineNumber}</span>
|
||||
<span className="text-gray-600 font-medium">{frame.fileName}</span>
|
||||
<span className="text-gray-500 mx-1">in</span>
|
||||
<span className="text-indigo-600 font-medium">{frame.inFunction || frame.fileName}</span>
|
||||
</div>
|
||||
<svg
|
||||
className={`w-5 h-5 text-gray-500 transition-transform ${expandedFrames[index] ? "transform rotate-180" : ""}`}
|
||||
fill="none"
|
||||
viewBox="0 0 24 24"
|
||||
stroke="currentColor"
|
||||
>
|
||||
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M19 9l-7 7-7-7" />
|
||||
</svg>
|
||||
</div>
|
||||
{(expandedFrames[index] || false) && frame.code && (
|
||||
<div className="px-12 py-2 font-mono text-sm text-gray-800 bg-gray-50 overflow-x-auto border-t border-gray-100">
|
||||
{frame.code}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
|
@ -1,267 +0,0 @@
|
|||
import { render, screen, act } from "@testing-library/react";
|
||||
import { describe, expect, it, vi, beforeEach } from "vitest";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { RequestResponsePanel } from "./RequestResponsePanel";
|
||||
import type { LogEntry } from "./columns";
|
||||
import NotificationsManager from "../molecules/notifications_manager";
|
||||
|
||||
const mockNotificationsManager = vi.mocked(NotificationsManager);
|
||||
|
||||
const baseLogEntry: LogEntry = {
|
||||
request_id: "chatcmpl-test-id",
|
||||
api_key: "api-key",
|
||||
team_id: "team-id",
|
||||
model: "gpt-4",
|
||||
model_id: "gpt-4",
|
||||
call_type: "chat",
|
||||
spend: 0,
|
||||
total_tokens: 0,
|
||||
prompt_tokens: 0,
|
||||
completion_tokens: 0,
|
||||
startTime: "2025-11-14T00:00:00Z",
|
||||
endTime: "2025-11-14T00:00:00Z",
|
||||
cache_hit: "miss",
|
||||
request_duration_ms: 1000,
|
||||
messages: [{ role: "user", content: "hello" }],
|
||||
response: { status: "ok" },
|
||||
metadata: {
|
||||
status: "success",
|
||||
additional_usage_values: {
|
||||
cache_read_input_tokens: 0,
|
||||
cache_creation_input_tokens: 0,
|
||||
},
|
||||
},
|
||||
request_tags: {},
|
||||
custom_llm_provider: "openai",
|
||||
api_base: "https://api.example.com",
|
||||
};
|
||||
|
||||
describe("RequestResponsePanel", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
Object.defineProperty(navigator, "clipboard", {
|
||||
value: {
|
||||
writeText: vi.fn().mockResolvedValue(undefined),
|
||||
},
|
||||
writable: true,
|
||||
configurable: true,
|
||||
});
|
||||
Object.defineProperty(window, "isSecureContext", {
|
||||
value: true,
|
||||
writable: true,
|
||||
configurable: true,
|
||||
});
|
||||
});
|
||||
|
||||
it("should render the component with request and response panels", () => {
|
||||
const mockGetRawRequest = vi.fn().mockReturnValue({ test: "request" });
|
||||
const mockFormattedResponse = vi.fn().mockReturnValue({ test: "response" });
|
||||
|
||||
render(
|
||||
<RequestResponsePanel
|
||||
row={{ original: baseLogEntry }}
|
||||
hasMessages={true}
|
||||
hasResponse={true}
|
||||
hasError={false}
|
||||
errorInfo={null}
|
||||
getRawRequest={mockGetRawRequest}
|
||||
formattedResponse={mockFormattedResponse}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(screen.getByText("Request")).toBeInTheDocument();
|
||||
expect(screen.getByText("Response")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should copy request to clipboard when copy button is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
const mockGetRawRequest = vi.fn().mockReturnValue({ test: "request data" });
|
||||
const mockFormattedResponse = vi.fn().mockReturnValue({ test: "response" });
|
||||
const mockWriteText = vi.fn().mockResolvedValue(undefined);
|
||||
|
||||
if (navigator.clipboard) {
|
||||
vi.spyOn(navigator.clipboard, "writeText").mockImplementation(mockWriteText);
|
||||
} else {
|
||||
Object.defineProperty(navigator, "clipboard", {
|
||||
value: {
|
||||
writeText: mockWriteText,
|
||||
},
|
||||
writable: true,
|
||||
configurable: true,
|
||||
});
|
||||
}
|
||||
|
||||
render(
|
||||
<RequestResponsePanel
|
||||
row={{ original: baseLogEntry }}
|
||||
hasMessages={true}
|
||||
hasResponse={true}
|
||||
hasError={false}
|
||||
errorInfo={null}
|
||||
getRawRequest={mockGetRawRequest}
|
||||
formattedResponse={mockFormattedResponse}
|
||||
/>,
|
||||
);
|
||||
|
||||
const copyButtons = screen.getAllByRole("button");
|
||||
const copyRequestButton = copyButtons.find((button) => button.getAttribute("title") === "Copy request");
|
||||
|
||||
expect(copyRequestButton).toBeInTheDocument();
|
||||
|
||||
await act(async () => {
|
||||
await user.click(copyRequestButton!);
|
||||
});
|
||||
|
||||
expect(mockGetRawRequest).toHaveBeenCalled();
|
||||
expect(mockWriteText).toHaveBeenCalledWith(JSON.stringify({ test: "request data" }, null, 2));
|
||||
expect(mockNotificationsManager.success).toHaveBeenCalledWith("Request copied to clipboard");
|
||||
});
|
||||
|
||||
it("should copy response to clipboard when copy button is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
const mockGetRawRequest = vi.fn().mockReturnValue({ test: "request" });
|
||||
const mockFormattedResponse = vi.fn().mockReturnValue({ test: "response data" });
|
||||
const mockWriteText = vi.fn().mockResolvedValue(undefined);
|
||||
|
||||
if (navigator.clipboard) {
|
||||
vi.spyOn(navigator.clipboard, "writeText").mockImplementation(mockWriteText);
|
||||
} else {
|
||||
Object.defineProperty(navigator, "clipboard", {
|
||||
value: {
|
||||
writeText: mockWriteText,
|
||||
},
|
||||
writable: true,
|
||||
configurable: true,
|
||||
});
|
||||
}
|
||||
|
||||
render(
|
||||
<RequestResponsePanel
|
||||
row={{ original: baseLogEntry }}
|
||||
hasMessages={true}
|
||||
hasResponse={true}
|
||||
hasError={false}
|
||||
errorInfo={null}
|
||||
getRawRequest={mockGetRawRequest}
|
||||
formattedResponse={mockFormattedResponse}
|
||||
/>,
|
||||
);
|
||||
|
||||
const copyButtons = screen.getAllByRole("button");
|
||||
const copyResponseButton = copyButtons.find((button) => button.getAttribute("title") === "Copy response");
|
||||
|
||||
expect(copyResponseButton).toBeInTheDocument();
|
||||
expect(copyResponseButton).not.toBeDisabled();
|
||||
|
||||
await act(async () => {
|
||||
await user.click(copyResponseButton!);
|
||||
});
|
||||
|
||||
expect(mockFormattedResponse).toHaveBeenCalled();
|
||||
expect(mockWriteText).toHaveBeenCalledWith(JSON.stringify({ test: "response data" }, null, 2));
|
||||
expect(mockNotificationsManager.success).toHaveBeenCalledWith("Response copied to clipboard");
|
||||
});
|
||||
|
||||
it("should call formattedResponse for the response panel and not getRawRequest", () => {
|
||||
const mockGetRawRequest = vi.fn().mockReturnValue({ requestData: "this should not appear in response" });
|
||||
const mockFormattedResponse = vi.fn().mockReturnValue({ responseData: "this should appear in response" });
|
||||
|
||||
render(
|
||||
<RequestResponsePanel
|
||||
row={{ original: baseLogEntry }}
|
||||
hasMessages={true}
|
||||
hasResponse={true}
|
||||
hasError={false}
|
||||
errorInfo={null}
|
||||
getRawRequest={mockGetRawRequest}
|
||||
formattedResponse={mockFormattedResponse}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(mockFormattedResponse).toHaveBeenCalled();
|
||||
expect(mockGetRawRequest).toHaveBeenCalled();
|
||||
|
||||
const formattedResponseCallCount = mockFormattedResponse.mock.calls.length;
|
||||
expect(formattedResponseCallCount).toBeGreaterThanOrEqual(1);
|
||||
|
||||
const responseData = mockFormattedResponse.mock.results[0].value;
|
||||
expect(responseData).toEqual({ responseData: "this should appear in response" });
|
||||
expect(responseData).not.toEqual({ requestData: "this should not appear in response" });
|
||||
});
|
||||
|
||||
it("should show error response data when hasError is true and hasResponse is false", () => {
|
||||
const failedLogEntry: LogEntry = {
|
||||
...baseLogEntry,
|
||||
messages: [],
|
||||
response: {},
|
||||
metadata: {
|
||||
status: "failure",
|
||||
error_information: {
|
||||
error_message: "Model not found",
|
||||
error_class: "NotFoundError",
|
||||
error_code: 404,
|
||||
},
|
||||
additional_usage_values: {
|
||||
cache_read_input_tokens: 0,
|
||||
cache_creation_input_tokens: 0,
|
||||
},
|
||||
},
|
||||
};
|
||||
const errorResponse = { error: { message: "Model not found", type: "NotFoundError", code: 404, param: null } };
|
||||
const mockGetRawRequest = vi.fn().mockReturnValue({ messages: [] });
|
||||
const mockFormattedResponse = vi.fn().mockReturnValue(errorResponse);
|
||||
render(
|
||||
<RequestResponsePanel
|
||||
row={{ original: failedLogEntry }}
|
||||
hasMessages={false}
|
||||
hasResponse={false}
|
||||
hasError={true}
|
||||
errorInfo={failedLogEntry.metadata.error_information}
|
||||
getRawRequest={mockGetRawRequest}
|
||||
formattedResponse={mockFormattedResponse}
|
||||
/>,
|
||||
);
|
||||
expect(screen.queryByText("Response data not available")).not.toBeInTheDocument();
|
||||
expect(mockFormattedResponse).toHaveBeenCalled();
|
||||
const copyButtons = screen.getAllByRole("button");
|
||||
const copyResponseButton = copyButtons.find((button) => button.getAttribute("title") === "Copy response");
|
||||
expect(copyResponseButton).not.toBeDisabled();
|
||||
});
|
||||
|
||||
it("should show Response data not available when hasResponse and hasError are both false", () => {
|
||||
const mockGetRawRequest = vi.fn().mockReturnValue({ messages: [] });
|
||||
const mockFormattedResponse = vi.fn().mockReturnValue({});
|
||||
render(
|
||||
<RequestResponsePanel
|
||||
row={{ original: baseLogEntry }}
|
||||
hasMessages={false}
|
||||
hasResponse={false}
|
||||
hasError={false}
|
||||
errorInfo={null}
|
||||
getRawRequest={mockGetRawRequest}
|
||||
formattedResponse={mockFormattedResponse}
|
||||
/>,
|
||||
);
|
||||
expect(screen.getByText("Response data not available")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show error code in response header when hasError is true", () => {
|
||||
const errorInfo = { error_message: "Rate limit exceeded", error_class: "RateLimitError", error_code: 429 };
|
||||
const mockGetRawRequest = vi.fn().mockReturnValue({ messages: [] });
|
||||
const mockFormattedResponse = vi
|
||||
.fn()
|
||||
.mockReturnValue({ error: { message: "Rate limit exceeded", type: "RateLimitError", code: 429, param: null } });
|
||||
render(
|
||||
<RequestResponsePanel
|
||||
row={{ original: baseLogEntry }}
|
||||
hasMessages={false}
|
||||
hasResponse={false}
|
||||
hasError={true}
|
||||
errorInfo={errorInfo}
|
||||
getRawRequest={mockGetRawRequest}
|
||||
formattedResponse={mockFormattedResponse}
|
||||
/>,
|
||||
);
|
||||
expect(screen.getByText(/HTTP code 429/)).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
@ -1,146 +0,0 @@
|
|||
import { LogEntry } from "./columns";
|
||||
import NotificationsManager from "../molecules/notifications_manager";
|
||||
import { JsonView, defaultStyles } from "react-json-view-lite";
|
||||
import "react-json-view-lite/dist/index.css";
|
||||
|
||||
interface RequestResponsePanelProps {
|
||||
row: {
|
||||
original: LogEntry;
|
||||
};
|
||||
hasMessages: string | boolean;
|
||||
hasResponse: string | boolean;
|
||||
hasError: boolean;
|
||||
errorInfo: any;
|
||||
getRawRequest: () => any;
|
||||
formattedResponse: () => any;
|
||||
}
|
||||
|
||||
export function RequestResponsePanel({
|
||||
row,
|
||||
hasMessages,
|
||||
hasResponse,
|
||||
hasError,
|
||||
errorInfo,
|
||||
getRawRequest,
|
||||
formattedResponse,
|
||||
}: RequestResponsePanelProps) {
|
||||
const copyToClipboard = async (text: string) => {
|
||||
try {
|
||||
// Try modern clipboard API first
|
||||
if (navigator.clipboard && window.isSecureContext) {
|
||||
await navigator.clipboard.writeText(text);
|
||||
return true;
|
||||
} else {
|
||||
// Fallback for non-secure contexts (like 0.0.0.0)
|
||||
const textArea = document.createElement("textarea");
|
||||
textArea.value = text;
|
||||
textArea.style.position = "fixed";
|
||||
textArea.style.opacity = "0";
|
||||
document.body.appendChild(textArea);
|
||||
textArea.focus();
|
||||
textArea.select();
|
||||
|
||||
const successful = document.execCommand("copy");
|
||||
document.body.removeChild(textArea);
|
||||
|
||||
if (!successful) {
|
||||
throw new Error("execCommand failed");
|
||||
}
|
||||
return true;
|
||||
}
|
||||
} catch (error) {
|
||||
console.error("Copy failed:", error);
|
||||
return false;
|
||||
}
|
||||
};
|
||||
|
||||
const handleCopyRequest = async () => {
|
||||
const success = await copyToClipboard(JSON.stringify(getRawRequest(), null, 2));
|
||||
if (success) {
|
||||
NotificationsManager.success("Request copied to clipboard");
|
||||
} else {
|
||||
NotificationsManager.fromBackend("Failed to copy request");
|
||||
}
|
||||
};
|
||||
|
||||
const handleCopyResponse = async () => {
|
||||
const success = await copyToClipboard(JSON.stringify(formattedResponse(), null, 2));
|
||||
if (success) {
|
||||
NotificationsManager.success("Response copied to clipboard");
|
||||
} else {
|
||||
NotificationsManager.fromBackend("Failed to copy response");
|
||||
}
|
||||
};
|
||||
|
||||
return (
|
||||
<div className="grid grid-cols-1 lg:grid-cols-2 gap-4 w-full max-w-full overflow-hidden box-border">
|
||||
{/* Request Side */}
|
||||
<div className="bg-white rounded-lg shadow w-full max-w-full overflow-hidden">
|
||||
<div className="flex justify-between items-center p-4 border-b">
|
||||
<h3 className="text-lg font-medium">Request</h3>
|
||||
<button onClick={handleCopyRequest} className="p-1 hover:bg-gray-200 rounded" title="Copy request">
|
||||
<svg
|
||||
xmlns="http://www.w3.org/2000/svg"
|
||||
width="16"
|
||||
height="16"
|
||||
viewBox="0 0 24 24"
|
||||
fill="none"
|
||||
stroke="currentColor"
|
||||
strokeWidth="2"
|
||||
strokeLinecap="round"
|
||||
strokeLinejoin="round"
|
||||
>
|
||||
<rect x="9" y="9" width="13" height="13" rx="2" ry="2"></rect>
|
||||
<path d="M5 15H4a2 2 0 0 1-2-2V4a2 2 0 0 1 2-2h9a2 2 0 0 1 2 2v1"></path>
|
||||
</svg>
|
||||
</button>
|
||||
</div>
|
||||
<div className="p-4 overflow-auto max-h-96 w-full max-w-full box-border">
|
||||
<div className="[&_[role='tree']]:bg-white [&_[role='tree']]:text-slate-900">
|
||||
<JsonView data={getRawRequest()} style={defaultStyles} clickToExpandNode={true} />
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* Response Side */}
|
||||
<div className="bg-white rounded-lg shadow w-full max-w-full overflow-hidden">
|
||||
<div className="flex justify-between items-center p-4 border-b">
|
||||
<h3 className="text-lg font-medium">
|
||||
Response
|
||||
{hasError && <span className="ml-2 text-sm text-red-600">• HTTP code {errorInfo?.error_code || 400}</span>}
|
||||
</h3>
|
||||
<button
|
||||
onClick={handleCopyResponse}
|
||||
className="p-1 hover:bg-gray-200 rounded"
|
||||
title="Copy response"
|
||||
disabled={!hasResponse && !hasError}
|
||||
>
|
||||
<svg
|
||||
xmlns="http://www.w3.org/2000/svg"
|
||||
width="16"
|
||||
height="16"
|
||||
viewBox="0 0 24 24"
|
||||
fill="none"
|
||||
stroke="currentColor"
|
||||
strokeWidth="2"
|
||||
strokeLinecap="round"
|
||||
strokeLinejoin="round"
|
||||
>
|
||||
<rect x="9" y="9" width="13" height="13" rx="2" ry="2"></rect>
|
||||
<path d="M5 15H4a2 2 0 0 1-2-2V4a2 2 0 0 1 2-2h9a2 2 0 0 1 2 2v1"></path>
|
||||
</svg>
|
||||
</button>
|
||||
</div>
|
||||
<div className="p-4 overflow-auto max-h-96 w-full max-w-full box-border">
|
||||
{hasResponse || hasError ? (
|
||||
<div className="[&_[role='tree']]:bg-white [&_[role='tree']]:text-slate-900">
|
||||
<JsonView data={formattedResponse()} style={defaultStyles} clickToExpandNode />
|
||||
</div>
|
||||
) : (
|
||||
<div className="text-gray-500 text-sm italic text-center py-4">Response data not available</div>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
|
@ -133,6 +133,13 @@ describe("migratedHref / legacyPageHref", () => {
|
|||
|
||||
expect(MIGRATED_PAGES.teams).toBe("teams");
|
||||
});
|
||||
|
||||
it("maps the organizations id to its route", async () => {
|
||||
vi.doMock("@/components/networking", () => ({ serverRootPath: "/" }));
|
||||
const { MIGRATED_PAGES } = await import("./migratedPages");
|
||||
|
||||
expect(MIGRATED_PAGES.organizations).toBe("organizations");
|
||||
});
|
||||
});
|
||||
|
||||
describe("dev server (NODE_ENV=development)", () => {
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@ export const MIGRATED_PAGES: Record<string, string> = {
|
|||
"router-settings": "router-settings",
|
||||
users: "users",
|
||||
teams: "teams",
|
||||
organizations: "organizations",
|
||||
};
|
||||
|
||||
function uiBase(): string {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue