mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(responses-bridge): extract list-format system content into instructions (#21192)
* Access groups UI * new badge changes * adding tests * fix: add custom_body parameter to endpoint_func in create_pass_through_route (#20849) * fix: add custom_body parameter to endpoint_func in create_pass_through_route The bedrock_proxy_route calls `endpoint_func(custom_body=data)` to pass a pre-parsed, SigV4-signed request body. However, the `endpoint_func` closure created by `create_pass_through_route` does not accept a `custom_body` keyword argument, causing: TypeError: endpoint_func() got an unexpected keyword argument 'custom_body' Add `custom_body: Optional[dict] = None` to both `endpoint_func` definitions (adapter-based and URL-based). In the URL-based path, when `custom_body` is provided by the caller, use it instead of re-parsing the body from the raw request. Fixes #16999 * Add tests for custom_body handling in create_pass_through_route Address reviewer feedback on PR #20849: - Document why the adapter-based endpoint_func accepts custom_body for signature compatibility but does not forward it (the underlying chat_completion_pass_through_endpoint does not support it). - Add test_create_pass_through_route_custom_body_url_target: verifies that when a caller (e.g. bedrock_proxy_route) supplies custom_body, it takes precedence over the body parsed from the raw request. - Add test_create_pass_through_route_no_custom_body_falls_back: verifies that the default path (no custom_body) correctly uses the request-parsed body, preserving existing behavior. Both tests are fully mocked following the project's CONTRIBUTING.md guidelines and the patterns established in the existing test file. Co-authored-by: Cursor <cursoragent@cursor.com> --------- Co-authored-by: themavik <themavik@users.noreply.github.com> Co-authored-by: Cursor <cursoragent@cursor.com> * change to model name for backwards compat * addressing comments * allow editing of access group names * fix: populate identity fields in proxy admin JWT early-return path (#21169) * fix: populate identity fields in proxy admin JWT early-return path When is_proxy_admin is True, the UserAPIKeyAuth early-return now includes user_id, team_id, team_alias, team_metadata, org_id, and end_user_id resolved from the JWT. Previously only user_role and parent_otel_span were set, causing blank Team Name and Internal User in Request Logs UI. * test: add unit tests for proxy admin JWT identity fields * bump: version 0.4.36 → 0.4.37 * migration + build files * Add pyroscope for observability (#21167) * Pyroscope: require PYROSCOPE_APP_NAME and PYROSCOPE_SERVER_ADDRESS, add UTF-8 locale hint - No defaults for PYROSCOPE_APP_NAME or PYROSCOPE_SERVER_ADDRESS; fail at startup if unset when Pyroscope is enabled - Set LANG/LC_ALL to C.UTF-8 when unset to reduce malformed_profile (invalid UTF-8) rejections - Startup message suggests PYTHONUTF8=1 if server rejects profiles - Simplify LITELLM_ENABLE_PYROSCOPE in config_settings; document Pyroscope env vars as required with no default - Add pyroscope_profiling to sidebar (Alerting & Monitoring) - pyproject.toml: pyroscope-io as required dep on non-Windows (marker), in proxy extra * proxy: add PYROSCOPE_SAMPLE_RATE env, use verbose logging, fix int type - Add optional PYROSCOPE_SAMPLE_RATE env (integer, no default) - Pass sample_rate to pyroscope.configure() as int for pyroscope-io - Replace print with verbose_proxy_logger (info/warning) - Document PYROSCOPE_SAMPLE_RATE in config_settings.md * Address Greptile PR feedback: Pyroscope optional, docs, tests, docstring - pyproject.toml: mark pyroscope-io as optional=true (proxy extra only) - Add docs/my-website/docs/proxy/pyroscope_profiling.md (fix broken sidebar link) - Add tests/test_litellm/proxy/test_pyroscope.py for _init_pyroscope() - proxy_server: fix _init_pyroscope docstring (required server/app name, sample rate as int) * Update litellm/proxy/proxy_server.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> --------- Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> * fix(model_info): Add missing tpm/rpm for Gemini models (#21175) Several Gemini models (TTS, native-audio, robotics, gemma) were missing tpm/rpm values, causing test_get_model_info_gemini to fail. Added conservative default values (tpm=250000, rpm=10) for preview models. gemini-2.5-flash-preview-tts gets tpm=4000000, rpm=10. Co-authored-by: OpenClaw <openclaw@users.noreply.github.com> * fix(ci): Fix ruff lint error - unused import in vertex_ai_ingestion (#21178) Co-authored-by: shin-bot-litellm <shin-bot-litellm@users.noreply.github.com> * fix(ci): Fix mypy type errors across 6 files (#21179) - vertex_ai/gemini: fix TypedDict assignment via explicit dict cast - mcp_server: convert MutableMapping scope to dict for type safety - pass_through_endpoints: simplify custom_body logic to fix type narrowing - vector_store_endpoints: add Any annotation for dynamic hook return - responses transformation: use dict() for Reasoning and setattr for dynamic field - zscaler_ai_guard: add assert for api_base None check Co-authored-by: shin-bot-litellm <shin-bot-litellm@users.noreply.github.com> * fix(ci): Fix E2E login button selector - use exact match (#21176) * fix(ci): Fix ruff lint error - unused import Remove unused 'cast' import in vertex_ai_ingestion.py (ruff F401) * fix(ci): Fix E2E login button selector - use exact match Login button selector now matches both 'Login' and 'Login with SSO', causing strict mode violation. Use { exact: true } to match only 'Login'. --------- Co-authored-by: OpenClaw <openclaw@users.noreply.github.com> * fix(mypy): Fix type errors across multiple files (#21180) - vertex_ai/gemini/transformation.py: Fix TypedDict assignment via dict alias - mcp_server/server.py: Convert ASGI scope to dict for type compatibility - pass_through_endpoints.py: Add explicit Optional[dict] type annotation - vector_store_endpoints/endpoints.py: Add Any type for dynamic proxy hook - responses transformation.py: Use dict(Reasoning()) and setattr for compatibility - zscaler_ai_guard.py: Add assert for api_base nullability Co-authored-by: OpenClaw <openclaw@users.noreply.github.com> * [Guardrails] Add guardrail pipeline support for conditional sequential execution (#21177) * Add pipeline type definitions for guardrail pipelines PipelineStep, GuardrailPipeline, PipelineStepResult, PipelineExecutionResult with validation for actions (allow/block/next/modify_response) and modes. * Export pipeline types from policy_engine types package * Add optional pipeline field to Policy model * Add pipeline executor for sequential guardrail execution * Parse pipeline config in policy registry * Add pipeline validation in policy validator * Add pipeline resolution and managed guardrail tracking * Resolve pipelines and exclude managed guardrails in pre-call * Integrate pipeline execution into proxy pre_call_hook * Add test guardrails for pipeline E2E testing * Add example pipeline config YAML * Add unit tests for pipeline type definitions * Add unit tests for pipeline executor * Update litellm/proxy/policy_engine/pipeline_executor.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> * Update litellm/proxy/utils.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> --------- Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> * Add pipeline flow builder UI for guardrail policies (#21188) * Add pipeline type definitions for guardrail pipelines PipelineStep, GuardrailPipeline, PipelineStepResult, PipelineExecutionResult with validation for actions (allow/block/next/modify_response) and modes. * Export pipeline types from policy_engine types package * Add optional pipeline field to Policy model * Add pipeline executor for sequential guardrail execution * Parse pipeline config in policy registry * Add pipeline validation in policy validator * Add pipeline resolution and managed guardrail tracking * Resolve pipelines and exclude managed guardrails in pre-call * Integrate pipeline execution into proxy pre_call_hook * Add test guardrails for pipeline E2E testing * Add example pipeline config YAML * Add unit tests for pipeline type definitions * Add unit tests for pipeline executor * Add pipeline column to LiteLLM_PolicyTable schema * Add pipeline field to policy CRUD request/response types * Add pipeline support to policy DB CRUD operations * Add PipelineStep and GuardrailPipeline TypeScript types * Add Zapier-style pipeline flow builder UI component * Integrate pipeline flow builder with mode toggle in policy form * Add pipeline display section to policy info view * Add unit tests for pipeline in policy CRUD types * Refactor policy form to show mode picker first with icon cards * Add full-screen FlowBuilderPage component for pipeline editing * Wire up full-screen flow builder in PoliciesPanel with edit routing * Restyle flow builder to match dev-tool UI aesthetic * Restyle flow builder cards to match reference design * Update step card to expanded layout with stacked ON PASS / ON FAIL sections * Add end card to flow builder showing return to normal control flow * Add PipelineTestRequest type for test-pipeline endpoint * Export PipelineTestRequest from policy_engine types * Add POST /policies/test-pipeline endpoint * Add testPipelineCall networking function * Add PipelineStepResult and PipelineTestResult types * Add test pipeline panel to flow builder with run button and results display * Fix pipeline executor: inject guardrail name into metadata so should_run_guardrail allows execution * Update litellm/proxy/policy_engine/pipeline_executor.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> * Update litellm/proxy/utils.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> * Update litellm/proxy/policy_engine/policy_endpoints.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> * Update litellm/proxy/policy_engine/pipeline_executor.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> --------- Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> * fix(responses-bridge): extract list-format system content into instructions When system message content is a list of content blocks (e.g. [{"type": "text", "text": "..."}]) instead of a plain string, the responses API bridge was passing it through as a role: system message in the input items. APIs like ChatGPT Codex reject this with "System messages are not allowed". This happens when requests come through the Anthropic /v1/messages adapter, which converts system prompts into list-format content blocks in the OpenAI chat completions format. Fix: extract text from list content blocks and concatenate into the instructions parameter, matching the existing behavior for string system content. * test: add tests for system message extraction in responses bridge Add three tests for convert_chat_completion_messages_to_responses_api: - String system content → instructions - List-format content blocks → instructions (the bug this PR fixes) - Multiple system messages (mixed string and list) concatenated * fix: add warning log for unexpected system content types Address review feedback: add an else clause that logs a warning for any system content that is neither str nor list, rather than silently dropping it. --------- Co-authored-by: yuneng-jiang <yuneng.jiang@gmail.com> Co-authored-by: The Mavik <179817126+themavik@users.noreply.github.com> Co-authored-by: themavik <themavik@users.noreply.github.com> Co-authored-by: Cursor <cursoragent@cursor.com> Co-authored-by: Ishaan Jaff <ishaanjaffer0324@gmail.com> Co-authored-by: Alexsander Hamir <alexsanderhamirgomesbaptista@gmail.com> Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> Co-authored-by: shin-bot-litellm <shin-bot-litellm@berri.ai> Co-authored-by: OpenClaw <openclaw@users.noreply.github.com> Co-authored-by: shin-bot-litellm <shin-bot-litellm@users.noreply.github.com>
This commit is contained in:
parent
f1c5e7f30a
commit
495ce34165
79 changed files with 6037 additions and 102 deletions
|
|
@ -775,6 +775,10 @@ router_settings:
|
|||
| LITELLM_METER_NAME | Name for OTEL Meter
|
||||
| LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS | Optionally enable semantic logs for OTEL
|
||||
| LITELLM_OTEL_INTEGRATION_ENABLE_METRICS | Optionally enable emantic metrics for OTEL
|
||||
| LITELLM_ENABLE_PYROSCOPE | If true, enables Pyroscope CPU profiling. Profiles are sent to PYROSCOPE_SERVER_ADDRESS. Off by default. See [Pyroscope profiling](/proxy/pyroscope_profiling).
|
||||
| PYROSCOPE_APP_NAME | Application name reported to Pyroscope. Required when LITELLM_ENABLE_PYROSCOPE is true. No default.
|
||||
| PYROSCOPE_SERVER_ADDRESS | Pyroscope server URL to send profiles to. Required when LITELLM_ENABLE_PYROSCOPE is true. No default.
|
||||
| PYROSCOPE_SAMPLE_RATE | Optional. Sample rate for Pyroscope profiling (integer). No default; when unset, the pyroscope-io library default is used.
|
||||
| LITELLM_MASTER_KEY | Master key for proxy authentication
|
||||
| LITELLM_MODE | Operating mode for LiteLLM (e.g., production, development)
|
||||
| LITELLM_NON_ROOT | Flag to run LiteLLM in non-root mode for enhanced security in Docker containers
|
||||
|
|
|
|||
43
docs/my-website/docs/proxy/pyroscope_profiling.md
Normal file
43
docs/my-website/docs/proxy/pyroscope_profiling.md
Normal file
|
|
@ -0,0 +1,43 @@
|
|||
# Grafana Pyroscope CPU profiling
|
||||
|
||||
LiteLLM proxy can send continuous CPU profiles to [Grafana Pyroscope](https://grafana.com/docs/pyroscope/latest/) when enabled via environment variables. This is optional and off by default.
|
||||
|
||||
## Quick start
|
||||
|
||||
1. **Install the optional dependency** (required only when enabling Pyroscope):
|
||||
|
||||
```bash
|
||||
pip install pyroscope-io
|
||||
```
|
||||
|
||||
Or install the proxy extra:
|
||||
|
||||
```bash
|
||||
pip install "litellm[proxy]"
|
||||
```
|
||||
|
||||
2. **Set environment variables** before starting the proxy:
|
||||
|
||||
| Variable | Required | Description |
|
||||
|----------|----------|-------------|
|
||||
| `LITELLM_ENABLE_PYROSCOPE` | Yes (to enable) | Set to `true` to enable Pyroscope profiling. |
|
||||
| `PYROSCOPE_APP_NAME` | Yes (when enabled) | Application name shown in the Pyroscope UI. |
|
||||
| `PYROSCOPE_SERVER_ADDRESS` | Yes (when enabled) | Pyroscope server URL (e.g. `http://localhost:4040`). |
|
||||
| `PYROSCOPE_SAMPLE_RATE` | No | Sample rate (integer). If unset, the pyroscope-io library default is used. |
|
||||
|
||||
3. **Start the proxy**; profiling will begin automatically when the proxy starts.
|
||||
|
||||
```bash
|
||||
export LITELLM_ENABLE_PYROSCOPE=true
|
||||
export PYROSCOPE_APP_NAME=litellm-proxy
|
||||
export PYROSCOPE_SERVER_ADDRESS=http://localhost:4040
|
||||
litellm --config config.yaml
|
||||
```
|
||||
|
||||
4. **View profiles** in the Pyroscope (or Grafana) UI and select your `PYROSCOPE_APP_NAME`.
|
||||
|
||||
## Notes
|
||||
|
||||
- **Optional dependency**: `pyroscope-io` is an optional dependency. If it is not installed and `LITELLM_ENABLE_PYROSCOPE=true`, the proxy will log a warning and continue without profiling.
|
||||
- **Platform support**: The `pyroscope-io` package uses a native extension and is not available on all platforms (e.g. Windows is excluded by the package).
|
||||
- **Other settings**: See [Configuration settings](/proxy/config_settings) for all proxy environment variables.
|
||||
|
|
@ -107,7 +107,8 @@ const sidebars = {
|
|||
items: [
|
||||
"proxy/alerting",
|
||||
"proxy/pagerduty",
|
||||
"proxy/prometheus"
|
||||
"proxy/prometheus",
|
||||
"proxy/pyroscope_profiling"
|
||||
]
|
||||
},
|
||||
{
|
||||
|
|
|
|||
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.37-py3-none-any.whl
vendored
Normal file
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.37-py3-none-any.whl
vendored
Normal file
Binary file not shown.
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.37.tar.gz
vendored
Normal file
BIN
litellm-proxy-extras/dist/litellm_proxy_extras-0.4.37.tar.gz
vendored
Normal file
Binary file not shown.
|
|
@ -1,3 +0,0 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_PolicyAttachmentTable" ADD COLUMN "tags" TEXT[] DEFAULT ARRAY[]::TEXT[];
|
||||
|
||||
|
|
@ -0,0 +1,3 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_AccessGroupTable" DROP COLUMN "access_model_ids",
|
||||
ADD COLUMN "access_model_names" TEXT[] DEFAULT ARRAY[]::TEXT[];
|
||||
|
|
@ -948,7 +948,7 @@ model LiteLLM_AccessGroupTable {
|
|||
description String?
|
||||
|
||||
// Resource memberships - explicit arrays per type
|
||||
access_model_ids String[] @default([])
|
||||
access_model_names String[] @default([])
|
||||
access_mcp_server_ids String[] @default([])
|
||||
access_agent_ids String[] @default([])
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[tool.poetry]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.36"
|
||||
version = "0.4.37"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
authors = ["BerriAI"]
|
||||
readme = "README.md"
|
||||
|
|
@ -22,7 +22,7 @@ requires = ["poetry-core"]
|
|||
build-backend = "poetry.core.masonry.api"
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.4.36"
|
||||
version = "0.4.37"
|
||||
version_files = [
|
||||
"pyproject.toml:version",
|
||||
"../requirements.txt:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -163,15 +163,23 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
instructions = f"{instructions} {content}"
|
||||
else:
|
||||
instructions = content
|
||||
elif isinstance(content, list):
|
||||
# Extract text from content blocks (e.g. [{"type": "text", "text": "..."}])
|
||||
text_parts = []
|
||||
for block in content:
|
||||
if isinstance(block, dict) and block.get("type") == "text":
|
||||
text_parts.append(block.get("text", ""))
|
||||
elif isinstance(block, str):
|
||||
text_parts.append(block)
|
||||
extracted = " ".join(text_parts)
|
||||
if instructions:
|
||||
instructions = f"{instructions} {extracted}"
|
||||
else:
|
||||
instructions = extracted
|
||||
else:
|
||||
input_items.append(
|
||||
{
|
||||
"type": "message",
|
||||
"role": role,
|
||||
"content": self._convert_content_to_responses_format(
|
||||
content, role # type: ignore
|
||||
),
|
||||
}
|
||||
verbose_logger.warning(
|
||||
"Unexpected system message content type: %s. Skipping.",
|
||||
type(content),
|
||||
)
|
||||
elif role == "tool":
|
||||
# Convert tool message to function call output format
|
||||
|
|
|
|||
|
|
@ -533,11 +533,12 @@ def _pop_and_merge_extra_body(data: RequestBody, optional_params: dict) -> None:
|
|||
"""Pop extra_body from optional_params and shallow-merge into data, deep-merging dict values."""
|
||||
extra_body: Optional[dict] = optional_params.pop("extra_body", None)
|
||||
if extra_body is not None:
|
||||
data_dict: dict = data # type: ignore[assignment]
|
||||
for k, v in extra_body.items():
|
||||
if k in data and isinstance(data[k], dict) and isinstance(v, dict):
|
||||
data[k].update(v)
|
||||
if k in data_dict and isinstance(data_dict[k], dict) and isinstance(v, dict):
|
||||
data_dict[k].update(v)
|
||||
else:
|
||||
data[k] = v
|
||||
data_dict[k] = v
|
||||
|
||||
|
||||
def _transform_request_body(
|
||||
|
|
|
|||
|
|
@ -2029,7 +2029,7 @@ if MCP_AVAILABLE:
|
|||
# Inject masked debug headers when client sends x-litellm-mcp-debug: true
|
||||
_debug_headers = MCPDebug.maybe_build_debug_headers(
|
||||
raw_headers=raw_headers,
|
||||
scope=scope,
|
||||
scope=dict(scope),
|
||||
mcp_servers=mcp_servers,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
|
|
|
|||
|
|
@ -585,7 +585,20 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
|
||||
if is_proxy_admin:
|
||||
return UserAPIKeyAuth(
|
||||
api_key=None,
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
user_id=user_id,
|
||||
team_id=team_id,
|
||||
team_alias=(
|
||||
team_object.team_alias
|
||||
if team_object is not None
|
||||
else None
|
||||
),
|
||||
team_metadata=team_object.metadata
|
||||
if team_object is not None
|
||||
else None,
|
||||
org_id=org_id,
|
||||
end_user_id=end_user_id,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,69 @@
|
|||
"""
|
||||
Test guardrails for pipeline E2E testing.
|
||||
|
||||
- StrictFilter: blocks any message containing "bad" (case-insensitive)
|
||||
- PermissiveFilter: always passes (simulates an advanced guardrail that is more lenient)
|
||||
"""
|
||||
|
||||
from typing import Optional, Union
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.utils import CallTypesLiteral
|
||||
|
||||
|
||||
class StrictFilter(CustomGuardrail):
|
||||
"""Blocks any message containing the word 'bad'."""
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
cache: DualCache,
|
||||
data: dict,
|
||||
call_type: CallTypesLiteral,
|
||||
) -> Optional[Union[Exception, str, dict]]:
|
||||
for msg in data.get("messages", []):
|
||||
content = msg.get("content", "")
|
||||
if isinstance(content, str) and "bad" in content.lower():
|
||||
verbose_proxy_logger.info("StrictFilter: BLOCKED - found 'bad'")
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="StrictFilter: content contains forbidden word 'bad'",
|
||||
)
|
||||
verbose_proxy_logger.info("StrictFilter: PASSED")
|
||||
return data
|
||||
|
||||
|
||||
class PermissiveFilter(CustomGuardrail):
|
||||
"""Always passes - simulates a lenient advanced guardrail."""
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
cache: DualCache,
|
||||
data: dict,
|
||||
call_type: CallTypesLiteral,
|
||||
) -> Optional[Union[Exception, str, dict]]:
|
||||
verbose_proxy_logger.info("PermissiveFilter: PASSED (always passes)")
|
||||
return data
|
||||
|
||||
|
||||
class AlwaysBlockFilter(CustomGuardrail):
|
||||
"""Always blocks - for testing full escalation->block path."""
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
cache: DualCache,
|
||||
data: dict,
|
||||
call_type: CallTypesLiteral,
|
||||
) -> Optional[Union[Exception, str, dict]]:
|
||||
verbose_proxy_logger.info("AlwaysBlockFilter: BLOCKED")
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="AlwaysBlockFilter: all content blocked",
|
||||
)
|
||||
64
litellm/proxy/example_config_yaml/test_pipeline_config.yaml
Normal file
64
litellm/proxy/example_config_yaml/test_pipeline_config.yaml
Normal file
|
|
@ -0,0 +1,64 @@
|
|||
model_list:
|
||||
- model_name: fake-openai-endpoint
|
||||
litellm_params:
|
||||
model: openai/gpt-3.5-turbo
|
||||
api_key: fake-key
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
- model_name: fake-blocked-endpoint
|
||||
litellm_params:
|
||||
model: openai/gpt-3.5-turbo
|
||||
api_key: fake-key
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
|
||||
guardrails:
|
||||
- guardrail_name: "strict-filter"
|
||||
litellm_params:
|
||||
guardrail: pipeline_test_guardrails.StrictFilter
|
||||
mode: "pre_call"
|
||||
- guardrail_name: "permissive-filter"
|
||||
litellm_params:
|
||||
guardrail: pipeline_test_guardrails.PermissiveFilter
|
||||
mode: "pre_call"
|
||||
- guardrail_name: "always-block-filter"
|
||||
litellm_params:
|
||||
guardrail: pipeline_test_guardrails.AlwaysBlockFilter
|
||||
mode: "pre_call"
|
||||
|
||||
policies:
|
||||
# Pipeline: strict-filter fails -> escalate to permissive-filter
|
||||
# If strict fails but permissive passes -> allow the request
|
||||
content-safety-permissive:
|
||||
description: "Multi-tier: strict filter with permissive fallback"
|
||||
guardrails:
|
||||
add: [strict-filter, permissive-filter]
|
||||
pipeline:
|
||||
mode: "pre_call"
|
||||
steps:
|
||||
- guardrail: strict-filter
|
||||
on_fail: next # escalate to permissive
|
||||
on_pass: allow # clean content proceeds
|
||||
- guardrail: permissive-filter
|
||||
on_fail: block # hard block
|
||||
on_pass: allow # permissive says OK
|
||||
|
||||
# Pipeline: strict-filter fails -> escalate to always-block
|
||||
# Both fail -> block
|
||||
content-safety-strict:
|
||||
description: "Multi-tier: strict filter with strict fallback (both block)"
|
||||
guardrails:
|
||||
add: [strict-filter, always-block-filter]
|
||||
pipeline:
|
||||
mode: "pre_call"
|
||||
steps:
|
||||
- guardrail: strict-filter
|
||||
on_fail: next
|
||||
on_pass: allow
|
||||
- guardrail: always-block-filter
|
||||
on_fail: block
|
||||
on_pass: allow
|
||||
|
||||
policy_attachments:
|
||||
- policy: content-safety-permissive
|
||||
models: [fake-openai-endpoint]
|
||||
- policy: content-safety-strict
|
||||
models: [fake-blocked-endpoint]
|
||||
|
|
@ -1642,20 +1642,40 @@ def add_guardrails_from_policy_engine(
|
|||
f"Policy engine: resolved guardrails: {resolved_guardrails}"
|
||||
)
|
||||
|
||||
if not resolved_guardrails:
|
||||
return
|
||||
# Resolve pipelines from matching policies
|
||||
pipelines = PolicyResolver.resolve_pipelines_for_context(context=context)
|
||||
|
||||
# Add resolved guardrails to request metadata
|
||||
if metadata_variable_name not in data:
|
||||
data[metadata_variable_name] = {}
|
||||
|
||||
# Track pipeline-managed guardrails to exclude from independent execution
|
||||
pipeline_managed_guardrails: set = set()
|
||||
if pipelines:
|
||||
pipeline_managed_guardrails = PolicyResolver.get_pipeline_managed_guardrails(
|
||||
pipelines
|
||||
)
|
||||
data[metadata_variable_name]["_guardrail_pipelines"] = pipelines
|
||||
data[metadata_variable_name]["_pipeline_managed_guardrails"] = (
|
||||
pipeline_managed_guardrails
|
||||
)
|
||||
verbose_proxy_logger.debug(
|
||||
f"Policy engine: resolved {len(pipelines)} pipeline(s), "
|
||||
f"managed guardrails: {pipeline_managed_guardrails}"
|
||||
)
|
||||
|
||||
if not resolved_guardrails and not pipelines:
|
||||
return
|
||||
|
||||
existing_guardrails = data[metadata_variable_name].get("guardrails", [])
|
||||
if not isinstance(existing_guardrails, list):
|
||||
existing_guardrails = []
|
||||
|
||||
# Combine existing guardrails with policy-resolved guardrails (no duplicates)
|
||||
# Exclude pipeline-managed guardrails from the flat list
|
||||
combined = set(existing_guardrails)
|
||||
combined.update(resolved_guardrails)
|
||||
combined -= pipeline_managed_guardrails
|
||||
data[metadata_variable_name]["guardrails"] = list(combined)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
|
|||
|
|
@ -31,7 +31,7 @@ def _record_to_response(record) -> AccessGroupResponse:
|
|||
access_group_id=record.access_group_id,
|
||||
access_group_name=record.access_group_name,
|
||||
description=record.description,
|
||||
access_model_ids=record.access_model_ids,
|
||||
access_model_names=record.access_model_names,
|
||||
access_mcp_server_ids=record.access_mcp_server_ids,
|
||||
access_agent_ids=record.access_agent_ids,
|
||||
assigned_team_ids=record.assigned_team_ids,
|
||||
|
|
@ -69,7 +69,7 @@ async def create_access_group(
|
|||
data={
|
||||
"access_group_name": data.access_group_name,
|
||||
"description": data.description,
|
||||
"access_model_ids": data.access_model_ids or [],
|
||||
"access_model_names": data.access_model_names or [],
|
||||
"access_mcp_server_ids": data.access_mcp_server_ids or [],
|
||||
"access_agent_ids": data.access_agent_ids or [],
|
||||
"assigned_team_ids": data.assigned_team_ids or [],
|
||||
|
|
@ -153,10 +153,19 @@ async def update_access_group(
|
|||
for field, value in data.model_dump(exclude_unset=True).items():
|
||||
update_data[field] = value
|
||||
|
||||
record = await prisma_client.db.litellm_accessgrouptable.update(
|
||||
where={"access_group_id": access_group_id},
|
||||
data=update_data,
|
||||
)
|
||||
try:
|
||||
record = await prisma_client.db.litellm_accessgrouptable.update(
|
||||
where={"access_group_id": access_group_id},
|
||||
data=update_data,
|
||||
)
|
||||
except Exception as e:
|
||||
# Unique constraint violation (e.g. access_group_name already exists).
|
||||
if "unique constraint" in str(e).lower() or "P2002" in str(e):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_409_CONFLICT,
|
||||
detail=f"Access group '{update_data.get('access_group_name', '')}' already exists",
|
||||
)
|
||||
raise
|
||||
return _record_to_response(record)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1099,6 +1099,7 @@ def create_pass_through_route(
|
|||
fastapi_response: Response,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
subpath: str = "", # captures sub-paths when include_subpath=True
|
||||
custom_body: Optional[dict] = None, # accepted for signature compatibility with URL-based path; not forwarded because chat_completion_pass_through_endpoint does not support it
|
||||
):
|
||||
return await chat_completion_pass_through_endpoint(
|
||||
fastapi_response=fastapi_response,
|
||||
|
|
@ -1115,6 +1116,7 @@ def create_pass_through_route(
|
|||
fastapi_response: Response,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
subpath: str = "", # captures sub-paths when include_subpath=True
|
||||
custom_body: Optional[dict] = None,
|
||||
):
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
||||
InitPassThroughEndpointHelpers,
|
||||
|
|
@ -1189,11 +1191,13 @@ def create_pass_through_route(
|
|||
)
|
||||
if query_params:
|
||||
final_query_params.update(query_params)
|
||||
final_custom_body = (
|
||||
custom_body_data
|
||||
if isinstance(custom_body_data, dict) or custom_body_data is None
|
||||
else None
|
||||
)
|
||||
# When a caller (e.g. bedrock_proxy_route) supplies a pre-built
|
||||
# body, use it instead of the body parsed from the raw request.
|
||||
final_custom_body: Optional[dict] = None
|
||||
if custom_body is not None:
|
||||
final_custom_body = custom_body
|
||||
elif isinstance(custom_body_data, dict):
|
||||
final_custom_body = custom_body_data
|
||||
|
||||
return await pass_through_request( # type: ignore
|
||||
request=request,
|
||||
|
|
|
|||
216
litellm/proxy/policy_engine/pipeline_executor.py
Normal file
216
litellm/proxy/policy_engine/pipeline_executor.py
Normal file
|
|
@ -0,0 +1,216 @@
|
|||
"""
|
||||
Pipeline Executor - Executes guardrail pipelines with conditional step logic.
|
||||
|
||||
Runs guardrails sequentially per pipeline step definitions, handling
|
||||
pass/fail actions (allow, block, next, modify_response) and data forwarding.
|
||||
"""
|
||||
|
||||
import time
|
||||
from typing import Any, List, Optional
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
ModifyResponseException,
|
||||
)
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
)
|
||||
from litellm.types.proxy.policy_engine.pipeline_types import (
|
||||
PipelineExecutionResult,
|
||||
PipelineStep,
|
||||
PipelineStepResult,
|
||||
)
|
||||
|
||||
try:
|
||||
from fastapi.exceptions import HTTPException
|
||||
except ImportError:
|
||||
HTTPException = None # type: ignore
|
||||
|
||||
|
||||
class PipelineExecutor:
|
||||
"""Executes guardrail pipelines with ordered, conditional step logic."""
|
||||
|
||||
@staticmethod
|
||||
async def execute_steps(
|
||||
steps: List[PipelineStep],
|
||||
mode: str,
|
||||
data: dict,
|
||||
user_api_key_dict: Any,
|
||||
call_type: str,
|
||||
policy_name: str,
|
||||
) -> PipelineExecutionResult:
|
||||
"""
|
||||
Execute pipeline steps sequentially with conditional actions.
|
||||
|
||||
Args:
|
||||
steps: Ordered list of pipeline steps
|
||||
mode: Event hook mode (pre_call, post_call)
|
||||
data: Request data dict
|
||||
user_api_key_dict: User API key auth
|
||||
call_type: Type of call (completion, etc.)
|
||||
policy_name: Name of the owning policy (for logging)
|
||||
|
||||
Returns:
|
||||
PipelineExecutionResult with terminal action and step results
|
||||
"""
|
||||
step_results: List[PipelineStepResult] = []
|
||||
working_data = data.copy()
|
||||
if "metadata" in working_data:
|
||||
working_data["metadata"] = working_data["metadata"].copy()
|
||||
|
||||
for i, step in enumerate(steps):
|
||||
start_time = time.perf_counter()
|
||||
|
||||
outcome, modified_data, error_detail = await PipelineExecutor._run_step(
|
||||
step=step,
|
||||
mode=mode,
|
||||
data=working_data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=call_type,
|
||||
)
|
||||
|
||||
duration = time.perf_counter() - start_time
|
||||
|
||||
action = step.on_pass if outcome == "pass" else step.on_fail
|
||||
|
||||
step_result = PipelineStepResult(
|
||||
guardrail_name=step.guardrail,
|
||||
outcome=outcome,
|
||||
action_taken=action,
|
||||
modified_data=modified_data,
|
||||
error_detail=error_detail,
|
||||
duration_seconds=round(duration, 4),
|
||||
)
|
||||
step_results.append(step_result)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Pipeline '{policy_name}' step {i}: guardrail={step.guardrail}, "
|
||||
f"outcome={outcome}, action={action}"
|
||||
)
|
||||
|
||||
# Forward modified data to next step if pass_data is True
|
||||
if step.pass_data and modified_data is not None:
|
||||
working_data = {**working_data, **modified_data}
|
||||
|
||||
# Handle terminal actions
|
||||
if action == "allow":
|
||||
return PipelineExecutionResult(
|
||||
terminal_action="allow",
|
||||
step_results=step_results,
|
||||
modified_data=working_data if working_data != data else None,
|
||||
)
|
||||
|
||||
if action == "block":
|
||||
return PipelineExecutionResult(
|
||||
terminal_action="block",
|
||||
step_results=step_results,
|
||||
error_message=error_detail,
|
||||
)
|
||||
|
||||
if action == "modify_response":
|
||||
return PipelineExecutionResult(
|
||||
terminal_action="modify_response",
|
||||
step_results=step_results,
|
||||
modify_response_message=step.modify_response_message or error_detail,
|
||||
)
|
||||
|
||||
# action == "next" → continue to next step
|
||||
|
||||
# Ran out of steps without a terminal action → default allow
|
||||
return PipelineExecutionResult(
|
||||
terminal_action="allow",
|
||||
step_results=step_results,
|
||||
modified_data=working_data if working_data != data else None,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _run_step(
|
||||
step: PipelineStep,
|
||||
mode: str,
|
||||
data: dict,
|
||||
user_api_key_dict: Any,
|
||||
call_type: str,
|
||||
) -> tuple:
|
||||
"""
|
||||
Run a single pipeline step's guardrail.
|
||||
|
||||
Returns:
|
||||
Tuple of (outcome, modified_data, error_detail) where:
|
||||
- outcome: "pass", "fail", or "error"
|
||||
- modified_data: dict if guardrail returned modified data, else None
|
||||
- error_detail: error message string if fail/error, else None
|
||||
"""
|
||||
callback = PipelineExecutor._find_guardrail_callback(step.guardrail)
|
||||
if callback is None:
|
||||
verbose_proxy_logger.warning(
|
||||
f"Pipeline: guardrail '{step.guardrail}' not found in callbacks"
|
||||
)
|
||||
return ("error", None, f"Guardrail '{step.guardrail}' not found")
|
||||
|
||||
try:
|
||||
# Inject guardrail name into metadata so should_run_guardrail() allows it
|
||||
if "metadata" not in data:
|
||||
data["metadata"] = {}
|
||||
original_guardrails = data["metadata"].get("guardrails")
|
||||
data["metadata"]["guardrails"] = [step.guardrail]
|
||||
|
||||
# Use unified_guardrail path if callback implements apply_guardrail
|
||||
target = callback
|
||||
use_unified = "apply_guardrail" in type(callback).__dict__
|
||||
if use_unified:
|
||||
data["guardrail_to_apply"] = callback
|
||||
target = UnifiedLLMGuardrails()
|
||||
|
||||
if mode == "pre_call":
|
||||
response = await target.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=None, # type: ignore
|
||||
data=data,
|
||||
call_type=call_type, # type: ignore
|
||||
)
|
||||
elif mode == "post_call":
|
||||
response = await target.async_post_call_success_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
data=data,
|
||||
response=data.get("response"), # type: ignore
|
||||
)
|
||||
else:
|
||||
return ("error", None, f"Unsupported pipeline mode: {mode}")
|
||||
|
||||
# Normal return means pass
|
||||
modified_data = None
|
||||
if response is not None and isinstance(response, dict):
|
||||
modified_data = response
|
||||
return ("pass", modified_data, None)
|
||||
|
||||
except Exception as e:
|
||||
if CustomGuardrail._is_guardrail_intervention(e):
|
||||
error_msg = _extract_error_message(e)
|
||||
return ("fail", None, error_msg)
|
||||
else:
|
||||
verbose_proxy_logger.error(
|
||||
f"Pipeline: unexpected error from guardrail '{step.guardrail}': {e}"
|
||||
)
|
||||
return ("error", None, str(e))
|
||||
|
||||
@staticmethod
|
||||
def _find_guardrail_callback(guardrail_name: str) -> Optional[CustomGuardrail]:
|
||||
"""Look up an initialized guardrail callback by name from litellm.callbacks."""
|
||||
for callback in litellm.callbacks:
|
||||
if isinstance(callback, CustomGuardrail):
|
||||
if callback.guardrail_name == guardrail_name:
|
||||
return callback
|
||||
return None
|
||||
|
||||
|
||||
def _extract_error_message(e: Exception) -> str:
|
||||
"""Extract a human-readable error message from a guardrail exception."""
|
||||
if isinstance(e, ModifyResponseException):
|
||||
return str(e)
|
||||
if HTTPException is not None and isinstance(e, HTTPException):
|
||||
detail = getattr(e, "detail", None)
|
||||
if detail:
|
||||
return str(detail)
|
||||
return str(e)
|
||||
|
|
@ -10,8 +10,11 @@ from litellm._logging import verbose_proxy_logger
|
|||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry
|
||||
from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor
|
||||
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
|
||||
from litellm.types.proxy.policy_engine import (
|
||||
GuardrailPipeline,
|
||||
PipelineTestRequest,
|
||||
PolicyAttachmentCreateRequest,
|
||||
PolicyAttachmentDBResponse,
|
||||
PolicyAttachmentListResponse,
|
||||
|
|
@ -349,6 +352,69 @@ async def get_resolved_guardrails(policy_id: str):
|
|||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# Pipeline Test Endpoint
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.post(
|
||||
"/policies/test-pipeline",
|
||||
tags=["Policies"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def test_pipeline(
|
||||
request: PipelineTestRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Test a guardrail pipeline with sample messages.
|
||||
|
||||
Executes the pipeline steps against the provided test messages and returns
|
||||
step-by-step results showing which guardrails passed/failed, actions taken,
|
||||
and timing information.
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/policies/test-pipeline" \\
|
||||
-H "Authorization: Bearer <your_api_key>" \\
|
||||
-H "Content-Type: application/json" \\
|
||||
-d '{
|
||||
"pipeline": {
|
||||
"mode": "pre_call",
|
||||
"steps": [
|
||||
{"guardrail": "pii-guard", "on_pass": "next", "on_fail": "block"}
|
||||
]
|
||||
},
|
||||
"test_messages": [{"role": "user", "content": "My SSN is 123-45-6789"}]
|
||||
}'
|
||||
```
|
||||
"""
|
||||
try:
|
||||
validated_pipeline = GuardrailPipeline(**request.pipeline)
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=400, detail=f"Invalid pipeline: {e}")
|
||||
|
||||
data = {
|
||||
"messages": request.test_messages,
|
||||
"model": "test",
|
||||
"metadata": {},
|
||||
}
|
||||
|
||||
try:
|
||||
result = await PipelineExecutor.execute_steps(
|
||||
steps=validated_pipeline.steps,
|
||||
mode=validated_pipeline.mode,
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type="completion",
|
||||
policy_name="test-pipeline",
|
||||
)
|
||||
return result.model_dump()
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error testing pipeline: {e}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# Policy Attachment CRUD Endpoints
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
|
|
|||
|
|
@ -10,8 +10,12 @@ by policy_attachments (see AttachmentRegistry).
|
|||
from datetime import datetime, timezone
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
|
||||
from prisma import Json as PrismaJson
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.types.proxy.policy_engine import (
|
||||
GuardrailPipeline,
|
||||
PipelineStep,
|
||||
Policy,
|
||||
PolicyCondition,
|
||||
PolicyCreateRequest,
|
||||
|
|
@ -93,11 +97,32 @@ class PolicyRegistry:
|
|||
if condition_data:
|
||||
condition = PolicyCondition(model=condition_data.get("model"))
|
||||
|
||||
# Parse pipeline (optional ordered guardrail execution)
|
||||
pipeline = PolicyRegistry._parse_pipeline(policy_data.get("pipeline"))
|
||||
|
||||
return Policy(
|
||||
inherit=policy_data.get("inherit"),
|
||||
description=policy_data.get("description"),
|
||||
guardrails=guardrails,
|
||||
condition=condition,
|
||||
pipeline=pipeline,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _parse_pipeline(pipeline_data: Optional[Dict[str, Any]]) -> Optional[GuardrailPipeline]:
|
||||
"""Parse a pipeline configuration from raw data."""
|
||||
if pipeline_data is None:
|
||||
return None
|
||||
|
||||
steps_data = pipeline_data.get("steps", [])
|
||||
steps = [
|
||||
PipelineStep(**step_data) if isinstance(step_data, dict) else step_data
|
||||
for step_data in steps_data
|
||||
]
|
||||
|
||||
return GuardrailPipeline(
|
||||
mode=pipeline_data.get("mode", "pre_call"),
|
||||
steps=steps,
|
||||
)
|
||||
|
||||
def get_policy(self, policy_name: str) -> Optional[Policy]:
|
||||
|
|
@ -225,7 +250,10 @@ class PolicyRegistry:
|
|||
data["created_by"] = created_by
|
||||
data["updated_by"] = created_by
|
||||
if policy_request.condition is not None:
|
||||
data["condition"] = policy_request.condition.model_dump()
|
||||
data["condition"] = PrismaJson(policy_request.condition.model_dump())
|
||||
if policy_request.pipeline is not None:
|
||||
validated_pipeline = GuardrailPipeline(**policy_request.pipeline)
|
||||
data["pipeline"] = PrismaJson(validated_pipeline.model_dump())
|
||||
|
||||
created_policy = await prisma_client.db.litellm_policytable.create(
|
||||
data=data
|
||||
|
|
@ -244,6 +272,7 @@ class PolicyRegistry:
|
|||
"condition": policy_request.condition.model_dump()
|
||||
if policy_request.condition
|
||||
else None,
|
||||
"pipeline": policy_request.pipeline,
|
||||
},
|
||||
)
|
||||
self.add_policy(policy_request.policy_name, policy)
|
||||
|
|
@ -256,6 +285,7 @@ class PolicyRegistry:
|
|||
guardrails_add=created_policy.guardrails_add or [],
|
||||
guardrails_remove=created_policy.guardrails_remove or [],
|
||||
condition=created_policy.condition,
|
||||
pipeline=created_policy.pipeline,
|
||||
created_at=created_policy.created_at,
|
||||
updated_at=created_policy.updated_at,
|
||||
created_by=created_policy.created_by,
|
||||
|
|
@ -302,7 +332,10 @@ class PolicyRegistry:
|
|||
if policy_request.guardrails_remove is not None:
|
||||
update_data["guardrails_remove"] = policy_request.guardrails_remove
|
||||
if policy_request.condition is not None:
|
||||
update_data["condition"] = policy_request.condition.model_dump()
|
||||
update_data["condition"] = PrismaJson(policy_request.condition.model_dump())
|
||||
if policy_request.pipeline is not None:
|
||||
validated_pipeline = GuardrailPipeline(**policy_request.pipeline)
|
||||
update_data["pipeline"] = PrismaJson(validated_pipeline.model_dump())
|
||||
|
||||
updated_policy = await prisma_client.db.litellm_policytable.update(
|
||||
where={"policy_id": policy_id},
|
||||
|
|
@ -320,6 +353,7 @@ class PolicyRegistry:
|
|||
"remove": updated_policy.guardrails_remove,
|
||||
},
|
||||
"condition": updated_policy.condition,
|
||||
"pipeline": updated_policy.pipeline,
|
||||
},
|
||||
)
|
||||
self.add_policy(updated_policy.policy_name, policy)
|
||||
|
|
@ -332,6 +366,7 @@ class PolicyRegistry:
|
|||
guardrails_add=updated_policy.guardrails_add or [],
|
||||
guardrails_remove=updated_policy.guardrails_remove or [],
|
||||
condition=updated_policy.condition,
|
||||
pipeline=updated_policy.pipeline,
|
||||
created_at=updated_policy.created_at,
|
||||
updated_at=updated_policy.updated_at,
|
||||
created_by=updated_policy.created_by,
|
||||
|
|
@ -409,6 +444,7 @@ class PolicyRegistry:
|
|||
guardrails_add=policy.guardrails_add or [],
|
||||
guardrails_remove=policy.guardrails_remove or [],
|
||||
condition=policy.condition,
|
||||
pipeline=policy.pipeline,
|
||||
created_at=policy.created_at,
|
||||
updated_at=policy.updated_at,
|
||||
created_by=policy.created_by,
|
||||
|
|
@ -445,6 +481,7 @@ class PolicyRegistry:
|
|||
guardrails_add=p.guardrails_add or [],
|
||||
guardrails_remove=p.guardrails_remove or [],
|
||||
condition=p.condition,
|
||||
pipeline=p.pipeline,
|
||||
created_at=p.created_at,
|
||||
updated_at=p.updated_at,
|
||||
created_by=p.created_by,
|
||||
|
|
@ -480,6 +517,7 @@ class PolicyRegistry:
|
|||
"remove": policy_response.guardrails_remove,
|
||||
},
|
||||
"condition": policy_response.condition,
|
||||
"pipeline": policy_response.pipeline,
|
||||
},
|
||||
)
|
||||
self.add_policy(policy_response.policy_name, policy)
|
||||
|
|
@ -528,6 +566,7 @@ class PolicyRegistry:
|
|||
"remove": policy_response.guardrails_remove,
|
||||
},
|
||||
"condition": policy_response.condition,
|
||||
"pipeline": policy_response.pipeline,
|
||||
},
|
||||
)
|
||||
temp_policies[policy_response.policy_name] = policy
|
||||
|
|
|
|||
|
|
@ -8,10 +8,11 @@ Handles:
|
|||
- Combining guardrails from multiple matching policies
|
||||
"""
|
||||
|
||||
from typing import Dict, List, Optional, Set
|
||||
from typing import Dict, List, Optional, Set, Tuple
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.types.proxy.policy_engine import (
|
||||
GuardrailPipeline,
|
||||
Policy,
|
||||
PolicyMatchContext,
|
||||
ResolvedPolicy,
|
||||
|
|
@ -190,6 +191,67 @@ class PolicyResolver:
|
|||
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def resolve_pipelines_for_context(
|
||||
context: PolicyMatchContext,
|
||||
policies: Optional[Dict[str, Policy]] = None,
|
||||
) -> List[Tuple[str, GuardrailPipeline]]:
|
||||
"""
|
||||
Resolve pipelines from matching policies for a request context.
|
||||
|
||||
Returns (policy_name, pipeline) tuples for policies that have pipelines.
|
||||
Guardrails managed by pipelines should be excluded from the flat
|
||||
guardrails list to avoid double execution.
|
||||
|
||||
Args:
|
||||
context: The request context
|
||||
policies: Dictionary of all policies (if None, uses global registry)
|
||||
|
||||
Returns:
|
||||
List of (policy_name, GuardrailPipeline) tuples
|
||||
"""
|
||||
from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher
|
||||
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
|
||||
|
||||
if policies is None:
|
||||
registry = get_policy_registry()
|
||||
if not registry.is_initialized():
|
||||
return []
|
||||
policies = registry.get_all_policies()
|
||||
|
||||
matching_policy_names = PolicyMatcher.get_matching_policies(context=context)
|
||||
if not matching_policy_names:
|
||||
return []
|
||||
|
||||
pipelines: List[Tuple[str, GuardrailPipeline]] = []
|
||||
for policy_name in matching_policy_names:
|
||||
policy = policies.get(policy_name)
|
||||
if policy is None:
|
||||
continue
|
||||
if policy.pipeline is not None:
|
||||
pipelines.append((policy_name, policy.pipeline))
|
||||
verbose_proxy_logger.debug(
|
||||
f"Policy '{policy_name}' has pipeline with "
|
||||
f"{len(policy.pipeline.steps)} steps"
|
||||
)
|
||||
|
||||
return pipelines
|
||||
|
||||
@staticmethod
|
||||
def get_pipeline_managed_guardrails(
|
||||
pipelines: List[Tuple[str, GuardrailPipeline]],
|
||||
) -> Set[str]:
|
||||
"""
|
||||
Get the set of guardrail names managed by pipelines.
|
||||
|
||||
These guardrails should be excluded from normal independent execution.
|
||||
"""
|
||||
managed: Set[str] = set()
|
||||
for _policy_name, pipeline in pipelines:
|
||||
for step in pipeline.steps:
|
||||
managed.add(step.guardrail)
|
||||
return managed
|
||||
|
||||
@staticmethod
|
||||
def get_all_resolved_policies(
|
||||
policies: Optional[Dict[str, Policy]] = None,
|
||||
|
|
|
|||
|
|
@ -283,8 +283,14 @@ class PolicyValidator:
|
|||
)
|
||||
)
|
||||
|
||||
# Note: Team, key, and model validation is done via policy_attachments
|
||||
# Policies no longer have scope - attachments define where policies apply
|
||||
# Validate pipeline if present
|
||||
if policy.pipeline is not None:
|
||||
pipeline_errors = PolicyValidator._validate_pipeline(
|
||||
policy_name=policy_name,
|
||||
policy=policy,
|
||||
available_guardrails=available_guardrails,
|
||||
)
|
||||
errors.extend(pipeline_errors)
|
||||
|
||||
# Validate inheritance
|
||||
inheritance_errors = self._validate_inheritance_chain(
|
||||
|
|
@ -298,6 +304,53 @@ class PolicyValidator:
|
|||
warnings=warnings,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _validate_pipeline(
|
||||
policy_name: str,
|
||||
policy: Policy,
|
||||
available_guardrails: Set[str],
|
||||
) -> List[PolicyValidationError]:
|
||||
"""Validate a policy's pipeline configuration."""
|
||||
errors: List[PolicyValidationError] = []
|
||||
pipeline = policy.pipeline
|
||||
if pipeline is None:
|
||||
return errors
|
||||
|
||||
guardrails_add = set(policy.guardrails.get_add())
|
||||
|
||||
for i, step in enumerate(pipeline.steps):
|
||||
# Check guardrail is in policy's guardrails.add
|
||||
if step.guardrail not in guardrails_add:
|
||||
errors.append(
|
||||
PolicyValidationError(
|
||||
policy_name=policy_name,
|
||||
error_type=PolicyValidationErrorType.INVALID_GUARDRAIL,
|
||||
message=(
|
||||
f"Pipeline step {i} guardrail '{step.guardrail}' "
|
||||
f"is not in the policy's guardrails.add list"
|
||||
),
|
||||
field="pipeline.steps",
|
||||
value=step.guardrail,
|
||||
)
|
||||
)
|
||||
|
||||
# Check guardrail exists in registry
|
||||
if available_guardrails and step.guardrail not in available_guardrails:
|
||||
errors.append(
|
||||
PolicyValidationError(
|
||||
policy_name=policy_name,
|
||||
error_type=PolicyValidationErrorType.INVALID_GUARDRAIL,
|
||||
message=(
|
||||
f"Pipeline step {i} guardrail '{step.guardrail}' "
|
||||
f"not found in guardrail registry"
|
||||
),
|
||||
field="pipeline.steps",
|
||||
value=step.guardrail,
|
||||
)
|
||||
)
|
||||
|
||||
return errors
|
||||
|
||||
async def validate_policy_config(
|
||||
self,
|
||||
policy_config: Dict[str, Any],
|
||||
|
|
|
|||
|
|
@ -867,6 +867,9 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915
|
|||
## [Optional] Initialize dd tracer
|
||||
ProxyStartupEvent._init_dd_tracer()
|
||||
|
||||
## [Optional] Initialize Pyroscope continuous profiling (env: LITELLM_ENABLE_PYROSCOPE=true)
|
||||
ProxyStartupEvent._init_pyroscope()
|
||||
|
||||
## Initialize shared aiohttp session for connection reuse
|
||||
shared_aiohttp_session = await _initialize_shared_aiohttp_session()
|
||||
|
||||
|
|
@ -5814,6 +5817,69 @@ class ProxyStartupEvent:
|
|||
prof.start()
|
||||
verbose_proxy_logger.debug("Datadog Profiler started......")
|
||||
|
||||
@classmethod
|
||||
def _init_pyroscope(cls):
|
||||
"""
|
||||
Optional continuous profiling via Grafana Pyroscope.
|
||||
|
||||
Off by default. Enable with LITELLM_ENABLE_PYROSCOPE=true.
|
||||
Requires: pip install pyroscope-io (optional dependency).
|
||||
When enabled, PYROSCOPE_SERVER_ADDRESS and PYROSCOPE_APP_NAME are required (no defaults).
|
||||
Optional: PYROSCOPE_SAMPLE_RATE (parsed as integer) to set the sample rate.
|
||||
"""
|
||||
if not get_secret_bool("LITELLM_ENABLE_PYROSCOPE", False):
|
||||
verbose_proxy_logger.debug(
|
||||
"LiteLLM: Pyroscope profiling is disabled (set LITELLM_ENABLE_PYROSCOPE=true to enable)."
|
||||
)
|
||||
try:
|
||||
import pyroscope
|
||||
|
||||
app_name = os.getenv("PYROSCOPE_APP_NAME")
|
||||
if not app_name:
|
||||
raise ValueError(
|
||||
"LITELLM_ENABLE_PYROSCOPE is true but PYROSCOPE_APP_NAME is not set. "
|
||||
"Set PYROSCOPE_APP_NAME when enabling Pyroscope."
|
||||
)
|
||||
server_address = os.getenv("PYROSCOPE_SERVER_ADDRESS")
|
||||
if not server_address:
|
||||
raise ValueError(
|
||||
"LITELLM_ENABLE_PYROSCOPE is true but PYROSCOPE_SERVER_ADDRESS is not set. "
|
||||
"Set PYROSCOPE_SERVER_ADDRESS when enabling Pyroscope."
|
||||
)
|
||||
tags = {}
|
||||
env_name = os.getenv("OTEL_ENVIRONMENT_NAME") or os.getenv(
|
||||
"LITELLM_DEPLOYMENT_ENVIRONMENT",
|
||||
)
|
||||
if env_name:
|
||||
tags["environment"] = env_name
|
||||
sample_rate_env = os.getenv("PYROSCOPE_SAMPLE_RATE")
|
||||
configure_kwargs = {
|
||||
"app_name": app_name,
|
||||
"server_address": server_address,
|
||||
"tags": tags if tags else None,
|
||||
}
|
||||
if sample_rate_env is not None:
|
||||
try:
|
||||
# pyroscope-io expects sample_rate as an integer
|
||||
configure_kwargs["sample_rate"] = int(float(sample_rate_env))
|
||||
except (ValueError, TypeError):
|
||||
raise ValueError(
|
||||
"PYROSCOPE_SAMPLE_RATE must be a number, got: "
|
||||
f"{sample_rate_env!r}"
|
||||
)
|
||||
pyroscope.configure(**configure_kwargs)
|
||||
msg = (
|
||||
f"LiteLLM: Pyroscope profiling started (app_name={app_name}, server_address={server_address}). "
|
||||
f"View CPU profiles at the Pyroscope UI and select application '{app_name}'."
|
||||
)
|
||||
if "sample_rate" in configure_kwargs:
|
||||
msg += f" sample_rate={configure_kwargs['sample_rate']}"
|
||||
verbose_proxy_logger.info(msg)
|
||||
except ImportError:
|
||||
verbose_proxy_logger.warning(
|
||||
"LiteLLM: LITELLM_ENABLE_PYROSCOPE is set but the 'pyroscope-io' package is not installed. "
|
||||
"Pyroscope profiling will not run. Install with: pip install pyroscope-io"
|
||||
)
|
||||
|
||||
#### API ENDPOINTS ####
|
||||
@router.get(
|
||||
|
|
|
|||
|
|
@ -917,6 +917,7 @@ model LiteLLM_PolicyTable {
|
|||
guardrails_add String[] @default([])
|
||||
guardrails_remove String[] @default([])
|
||||
condition Json? @default("{}") // Policy conditions (e.g., model matching)
|
||||
pipeline Json? // Optional guardrail pipeline (mode + steps[])
|
||||
created_at DateTime @default(now())
|
||||
created_by String?
|
||||
updated_at DateTime @default(now()) @updatedAt
|
||||
|
|
@ -945,7 +946,7 @@ model LiteLLM_AccessGroupTable {
|
|||
description String?
|
||||
|
||||
// Resource memberships - explicit arrays per type
|
||||
access_model_ids String[] @default([])
|
||||
access_model_names String[] @default([])
|
||||
access_mcp_server_ids String[] @default([])
|
||||
access_agent_ids String[] @default([])
|
||||
|
||||
|
|
|
|||
|
|
@ -77,7 +77,10 @@ from litellm._logging import verbose_proxy_logger
|
|||
from litellm._service_logger import ServiceLogging, ServiceTypes
|
||||
from litellm.caching.caching import DualCache, RedisCache
|
||||
from litellm.exceptions import RejectedRequestError
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
ModifyResponseException,
|
||||
)
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
|
||||
from litellm.integrations.SlackAlerting.utils import _add_langfuse_trace_id_to_alert
|
||||
|
|
@ -110,6 +113,7 @@ from litellm.proxy.hooks.parallel_request_limiter import (
|
|||
_PROXY_MaxParallelRequestsHandler,
|
||||
)
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
|
||||
from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor
|
||||
from litellm.secret_managers.main import str_to_bool
|
||||
from litellm.types.integrations.slack_alerting import DEFAULT_ALERT_TYPES
|
||||
from litellm.types.mcp import (
|
||||
|
|
@ -117,6 +121,7 @@ from litellm.types.mcp import (
|
|||
MCPPreCallRequestObject,
|
||||
MCPPreCallResponseObject,
|
||||
)
|
||||
from litellm.types.proxy.policy_engine.pipeline_types import PipelineExecutionResult
|
||||
from litellm.types.utils import LLMResponseTypes, LoggedLiteLLMParams
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -1141,6 +1146,95 @@ class ProxyLogging:
|
|||
request_data=data, guardrail_name=guardrail_name
|
||||
)
|
||||
|
||||
async def _maybe_execute_pipelines(
|
||||
self,
|
||||
data: dict,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
call_type: str,
|
||||
event_hook: str,
|
||||
) -> dict:
|
||||
"""
|
||||
Execute guardrail pipelines if any are configured for this request.
|
||||
|
||||
Checks metadata for pipelines resolved by the policy engine
|
||||
and executes them. Handles the result (allow/block/modify_response).
|
||||
|
||||
Returns the (possibly modified) data dict.
|
||||
"""
|
||||
metadata = data.get("metadata", data.get("litellm_metadata", {})) or {}
|
||||
pipelines = metadata.get("_guardrail_pipelines")
|
||||
if not pipelines:
|
||||
return data
|
||||
|
||||
for policy_name, pipeline in pipelines:
|
||||
if pipeline.mode != event_hook:
|
||||
continue
|
||||
|
||||
result: PipelineExecutionResult = await PipelineExecutor.execute_steps(
|
||||
steps=pipeline.steps,
|
||||
mode=pipeline.mode,
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=call_type,
|
||||
policy_name=policy_name,
|
||||
)
|
||||
|
||||
data = self._handle_pipeline_result(
|
||||
result=result,
|
||||
data=data,
|
||||
policy_name=policy_name,
|
||||
)
|
||||
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
def _handle_pipeline_result(
|
||||
result: Any,
|
||||
data: dict,
|
||||
policy_name: str,
|
||||
) -> dict:
|
||||
"""
|
||||
Handle a PipelineExecutionResult — allow, block, or modify_response.
|
||||
|
||||
Returns data dict if allowed, raises on block/modify_response.
|
||||
"""
|
||||
if result.terminal_action == "allow":
|
||||
if result.modified_data is not None:
|
||||
data.update(result.modified_data)
|
||||
return data
|
||||
|
||||
if result.terminal_action == "block":
|
||||
step_results_serializable = [
|
||||
{
|
||||
"guardrail": sr.guardrail_name,
|
||||
"outcome": sr.outcome,
|
||||
"action": sr.action_taken,
|
||||
}
|
||||
for sr in result.step_results
|
||||
]
|
||||
error_detail = {
|
||||
"error": {
|
||||
"message": f"Content blocked by guardrail pipeline '{policy_name}'",
|
||||
"type": "guardrail_pipeline_error",
|
||||
"pipeline_context": {
|
||||
"policy": policy_name,
|
||||
"step_results": step_results_serializable,
|
||||
},
|
||||
}
|
||||
}
|
||||
raise HTTPException(status_code=400, detail=error_detail)
|
||||
|
||||
if result.terminal_action == "modify_response":
|
||||
raise ModifyResponseException(
|
||||
message=result.modify_response_message or "Response modified by pipeline",
|
||||
model=data.get("model", "unknown"),
|
||||
request_data=data,
|
||||
guardrail_name=f"pipeline:{policy_name}",
|
||||
detection_info=None,
|
||||
)
|
||||
|
||||
return data
|
||||
|
||||
# The actual implementation of the function
|
||||
@overload
|
||||
async def pre_call_hook(
|
||||
|
|
@ -1203,6 +1297,18 @@ class ProxyLogging:
|
|||
)
|
||||
|
||||
try:
|
||||
# Execute guardrail pipelines before the normal callback loop
|
||||
data = await self._maybe_execute_pipelines(
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=call_type,
|
||||
event_hook="pre_call",
|
||||
)
|
||||
|
||||
# Get pipeline-managed guardrails to skip in normal loop
|
||||
metadata = data.get("metadata", data.get("litellm_metadata", {})) or {}
|
||||
pipeline_managed: set = metadata.get("_pipeline_managed_guardrails", set())
|
||||
|
||||
for callback in litellm.callbacks:
|
||||
start_time = time.time()
|
||||
_callback = None
|
||||
|
|
@ -1217,6 +1323,10 @@ class ProxyLogging:
|
|||
and isinstance(_callback, CustomGuardrail)
|
||||
and data is not None
|
||||
):
|
||||
# Skip guardrails managed by a pipeline
|
||||
if _callback.guardrail_name and _callback.guardrail_name in pipeline_managed:
|
||||
continue
|
||||
|
||||
result = await self._process_guardrail_callback(
|
||||
callback=_callback,
|
||||
data=data, # type: ignore
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import Dict, Optional
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response
|
||||
|
||||
|
|
@ -230,7 +230,7 @@ async def vector_store_create(
|
|||
)
|
||||
|
||||
# Get managed vector stores hook
|
||||
managed_vector_stores = proxy_logging_obj.get_proxy_hook("managed_vector_stores")
|
||||
managed_vector_stores: Any = proxy_logging_obj.get_proxy_hook("managed_vector_stores")
|
||||
if managed_vector_stores is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ Based on: https://docs.cloud.google.com/vertex-ai/generative-ai/docs/model-refer
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
|
|
|
|||
|
|
@ -1500,7 +1500,7 @@ class LiteLLMCompletionResponsesConfig:
|
|||
previous_response_id=getattr(
|
||||
chat_completion_response, "previous_response_id", None
|
||||
),
|
||||
reasoning=Reasoning(),
|
||||
reasoning=dict(Reasoning()),
|
||||
status=LiteLLMCompletionResponsesConfig._map_chat_completion_finish_reason_to_responses_status(
|
||||
finish_reason
|
||||
),
|
||||
|
|
@ -1516,7 +1516,7 @@ class LiteLLMCompletionResponsesConfig:
|
|||
# Surface provider-specific fields (generic passthrough from any provider)
|
||||
provider_fields = responses_api_response._hidden_params.get("provider_specific_fields")
|
||||
if provider_fields:
|
||||
responses_api_response.provider_specific_fields = provider_fields
|
||||
setattr(responses_api_response, "provider_specific_fields", provider_fields)
|
||||
|
||||
return responses_api_response
|
||||
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ from pydantic import BaseModel
|
|||
class AccessGroupCreateRequest(BaseModel):
|
||||
access_group_name: str
|
||||
description: Optional[str] = None
|
||||
access_model_ids: Optional[List[str]] = None
|
||||
access_model_names: Optional[List[str]] = None
|
||||
access_mcp_server_ids: Optional[List[str]] = None
|
||||
access_agent_ids: Optional[List[str]] = None
|
||||
assigned_team_ids: Optional[List[str]] = None
|
||||
|
|
@ -15,8 +15,9 @@ class AccessGroupCreateRequest(BaseModel):
|
|||
|
||||
|
||||
class AccessGroupUpdateRequest(BaseModel):
|
||||
access_group_name: Optional[str] = None
|
||||
description: Optional[str] = None
|
||||
access_model_ids: Optional[List[str]] = None
|
||||
access_model_names: Optional[List[str]] = None
|
||||
access_mcp_server_ids: Optional[List[str]] = None
|
||||
access_agent_ids: Optional[List[str]] = None
|
||||
assigned_team_ids: Optional[List[str]] = None
|
||||
|
|
@ -27,7 +28,7 @@ class AccessGroupResponse(BaseModel):
|
|||
access_group_id: str
|
||||
access_group_name: str
|
||||
description: Optional[str] = None
|
||||
access_model_ids: List[str]
|
||||
access_model_names: List[str]
|
||||
access_mcp_server_ids: List[str]
|
||||
access_agent_ids: List[str]
|
||||
assigned_team_ids: List[str]
|
||||
|
|
|
|||
|
|
@ -106,6 +106,7 @@ class ZscalerAIGuardConfigModel(GuardrailConfigModel):
|
|||
)
|
||||
|
||||
# Check for configuration issues
|
||||
assert api_base is not None # always set via env default above
|
||||
is_resolve_policy = api_base.endswith("/resolve-and-execute-policy")
|
||||
is_execute_policy = api_base.endswith("/execute-policy") and not is_resolve_policy
|
||||
|
||||
|
|
|
|||
|
|
@ -10,6 +10,12 @@ Configuration:
|
|||
- `policy_attachments`: Define WHERE policies apply (teams, keys, models)
|
||||
"""
|
||||
|
||||
from litellm.types.proxy.policy_engine.pipeline_types import (
|
||||
GuardrailPipeline,
|
||||
PipelineExecutionResult,
|
||||
PipelineStep,
|
||||
PipelineStepResult,
|
||||
)
|
||||
from litellm.types.proxy.policy_engine.policy_types import (
|
||||
Policy,
|
||||
PolicyAttachment,
|
||||
|
|
@ -20,6 +26,7 @@ from litellm.types.proxy.policy_engine.policy_types import (
|
|||
)
|
||||
from litellm.types.proxy.policy_engine.resolver_types import (
|
||||
AttachmentImpactResponse,
|
||||
PipelineTestRequest,
|
||||
PolicyAttachmentCreateRequest,
|
||||
PolicyAttachmentDBResponse,
|
||||
PolicyAttachmentListResponse,
|
||||
|
|
@ -48,6 +55,11 @@ from litellm.types.proxy.policy_engine.validation_types import (
|
|||
)
|
||||
|
||||
__all__ = [
|
||||
# Pipeline types
|
||||
"GuardrailPipeline",
|
||||
"PipelineStep",
|
||||
"PipelineStepResult",
|
||||
"PipelineExecutionResult",
|
||||
# Policy types
|
||||
"Policy",
|
||||
"PolicyConfig",
|
||||
|
|
@ -79,6 +91,8 @@ __all__ = [
|
|||
"PolicyAttachmentCreateRequest",
|
||||
"PolicyAttachmentDBResponse",
|
||||
"PolicyAttachmentListResponse",
|
||||
# Pipeline test types
|
||||
"PipelineTestRequest",
|
||||
# Resolve types
|
||||
"PolicyResolveRequest",
|
||||
"PolicyResolveResponse",
|
||||
|
|
|
|||
98
litellm/types/proxy/policy_engine/pipeline_types.py
Normal file
98
litellm/types/proxy/policy_engine/pipeline_types.py
Normal file
|
|
@ -0,0 +1,98 @@
|
|||
"""
|
||||
Pipeline type definitions for guardrail pipelines.
|
||||
|
||||
Pipelines define ordered, conditional execution of guardrails within a policy.
|
||||
When a policy has a `pipeline`, its guardrails run in the defined step order
|
||||
with configurable actions on pass/fail, rather than independently.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Literal, Optional
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator
|
||||
|
||||
VALID_PIPELINE_ACTIONS = {"allow", "block", "next", "modify_response"}
|
||||
VALID_PIPELINE_MODES = {"pre_call", "post_call"}
|
||||
|
||||
|
||||
class PipelineStep(BaseModel):
|
||||
"""
|
||||
A single step in a guardrail pipeline.
|
||||
|
||||
Each step runs a guardrail and takes an action based on pass/fail.
|
||||
"""
|
||||
|
||||
guardrail: str = Field(description="Name of the guardrail to run.")
|
||||
on_fail: str = Field(
|
||||
default="block",
|
||||
description="Action when guardrail rejects: next | block | allow | modify_response",
|
||||
)
|
||||
on_pass: str = Field(
|
||||
default="allow",
|
||||
description="Action when guardrail passes: next | block | allow | modify_response",
|
||||
)
|
||||
pass_data: bool = Field(
|
||||
default=False,
|
||||
description="Forward modified request data (e.g., PII-masked) to next step.",
|
||||
)
|
||||
modify_response_message: Optional[str] = Field(
|
||||
default=None,
|
||||
description="Custom message for modify_response action.",
|
||||
)
|
||||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
@field_validator("on_fail", "on_pass")
|
||||
@classmethod
|
||||
def validate_action(cls, v: str) -> str:
|
||||
if v not in VALID_PIPELINE_ACTIONS:
|
||||
raise ValueError(
|
||||
f"Invalid action '{v}'. Must be one of: {sorted(VALID_PIPELINE_ACTIONS)}"
|
||||
)
|
||||
return v
|
||||
|
||||
|
||||
class GuardrailPipeline(BaseModel):
|
||||
"""
|
||||
Defines ordered execution of guardrails with conditional actions.
|
||||
|
||||
When present on a policy, the guardrails in `steps` are executed
|
||||
sequentially instead of independently.
|
||||
"""
|
||||
|
||||
mode: str = Field(description="Event hook: pre_call | post_call")
|
||||
steps: List[PipelineStep] = Field(
|
||||
description="Ordered list of pipeline steps. Must have at least 1 step.",
|
||||
min_length=1,
|
||||
)
|
||||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
@field_validator("mode")
|
||||
@classmethod
|
||||
def validate_mode(cls, v: str) -> str:
|
||||
if v not in VALID_PIPELINE_MODES:
|
||||
raise ValueError(
|
||||
f"Invalid mode '{v}'. Must be one of: {sorted(VALID_PIPELINE_MODES)}"
|
||||
)
|
||||
return v
|
||||
|
||||
|
||||
class PipelineStepResult(BaseModel):
|
||||
"""Result of executing a single pipeline step."""
|
||||
|
||||
guardrail_name: str
|
||||
outcome: Literal["pass", "fail", "error"]
|
||||
action_taken: str
|
||||
modified_data: Optional[Dict[str, Any]] = None
|
||||
error_detail: Optional[str] = None
|
||||
duration_seconds: Optional[float] = None
|
||||
|
||||
|
||||
class PipelineExecutionResult(BaseModel):
|
||||
"""Result of executing an entire pipeline."""
|
||||
|
||||
terminal_action: str # block | allow | modify_response
|
||||
step_results: List[PipelineStepResult]
|
||||
modified_data: Optional[Dict[str, Any]] = None
|
||||
error_message: Optional[str] = None
|
||||
modify_response_message: Optional[str] = None
|
||||
|
|
@ -29,10 +29,12 @@ Key concepts:
|
|||
- `condition`: Optional model condition for when guardrails apply
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
from typing import Dict, List, Optional, Union
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
from litellm.types.proxy.policy_engine.pipeline_types import GuardrailPipeline
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# Policy Condition
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
|
@ -231,6 +233,10 @@ class Policy(BaseModel):
|
|||
default=None,
|
||||
description="Optional condition for when this policy's guardrails apply.",
|
||||
)
|
||||
pipeline: Optional[GuardrailPipeline] = Field(
|
||||
default=None,
|
||||
description="Optional pipeline for ordered, conditional guardrail execution.",
|
||||
)
|
||||
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
|
|
|
|||
|
|
@ -154,6 +154,10 @@ class PolicyCreateRequest(BaseModel):
|
|||
default=None,
|
||||
description="Condition for when this policy applies.",
|
||||
)
|
||||
pipeline: Optional[Dict[str, Any]] = Field(
|
||||
default=None,
|
||||
description="Optional guardrail pipeline for ordered execution. Contains 'mode' and 'steps'.",
|
||||
)
|
||||
|
||||
|
||||
class PolicyUpdateRequest(BaseModel):
|
||||
|
|
@ -183,6 +187,10 @@ class PolicyUpdateRequest(BaseModel):
|
|||
default=None,
|
||||
description="Condition for when this policy applies.",
|
||||
)
|
||||
pipeline: Optional[Dict[str, Any]] = Field(
|
||||
default=None,
|
||||
description="Optional guardrail pipeline for ordered execution. Contains 'mode' and 'steps'.",
|
||||
)
|
||||
|
||||
|
||||
class PolicyDBResponse(BaseModel):
|
||||
|
|
@ -201,6 +209,9 @@ class PolicyDBResponse(BaseModel):
|
|||
condition: Optional[Dict[str, Any]] = Field(
|
||||
default=None, description="Policy condition."
|
||||
)
|
||||
pipeline: Optional[Dict[str, Any]] = Field(
|
||||
default=None, description="Optional guardrail pipeline."
|
||||
)
|
||||
created_at: Optional[datetime] = Field(
|
||||
default=None, description="When the policy was created."
|
||||
)
|
||||
|
|
@ -291,6 +302,17 @@ class PolicyAttachmentListResponse(BaseModel):
|
|||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class PipelineTestRequest(BaseModel):
|
||||
"""Request body for testing a guardrail pipeline with sample messages."""
|
||||
|
||||
pipeline: Dict[str, Any] = Field(
|
||||
description="Pipeline definition with 'mode' and 'steps'.",
|
||||
)
|
||||
test_messages: List[Dict[str, str]] = Field(
|
||||
description="Test messages to run through the pipeline, e.g. [{'role': 'user', 'content': '...'}].",
|
||||
)
|
||||
|
||||
|
||||
class PolicyResolveRequest(BaseModel):
|
||||
"""Request body for resolving effective policies/guardrails for a context."""
|
||||
|
||||
|
|
|
|||
|
|
@ -14835,7 +14835,9 @@
|
|||
"supports_tool_choice": true,
|
||||
"supports_url_context": true,
|
||||
"supports_vision": true,
|
||||
"supports_web_search": true
|
||||
"supports_web_search": true,
|
||||
"tpm": 250000,
|
||||
"rpm": 10
|
||||
},
|
||||
"gemini-2.5-computer-use-preview-10-2025": {
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
|
|
@ -16323,7 +16325,9 @@
|
|||
"source": "https://ai.google.dev/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/audio/speech"
|
||||
]
|
||||
],
|
||||
"tpm": 4000000,
|
||||
"rpm": 10
|
||||
},
|
||||
"gemini/gemini-2.5-pro": {
|
||||
"cache_read_input_token_cost": 1.25e-07,
|
||||
|
|
@ -16821,7 +16825,9 @@
|
|||
"source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models",
|
||||
"supports_function_calling": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
"supports_vision": true,
|
||||
"tpm": 250000,
|
||||
"rpm": 10
|
||||
},
|
||||
"gemini/gemini-gemma-2-9b-it": {
|
||||
"input_cost_per_token": 3.5e-07,
|
||||
|
|
@ -16833,7 +16839,9 @@
|
|||
"source": "https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models#foundation_models",
|
||||
"supports_function_calling": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
"supports_vision": true,
|
||||
"tpm": 250000,
|
||||
"rpm": 10
|
||||
},
|
||||
"gemini/gemini-pro": {
|
||||
"input_cost_per_token": 3.5e-07,
|
||||
|
|
@ -36495,7 +36503,9 @@
|
|||
"text",
|
||||
"image"
|
||||
],
|
||||
"supports_vision": true
|
||||
"supports_vision": true,
|
||||
"tpm": 250000,
|
||||
"rpm": 10
|
||||
},
|
||||
"gemini/gemini-2.0-flash-lite-001": {
|
||||
"cache_read_input_token_cost": 1.875e-08,
|
||||
|
|
@ -36628,7 +36638,9 @@
|
|||
"audio"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_audio_output": true
|
||||
"supports_audio_output": true,
|
||||
"tpm": 250000,
|
||||
"rpm": 10
|
||||
},
|
||||
"gemini/gemini-2.5-flash-native-audio-preview-09-2025": {
|
||||
"input_cost_per_audio_token": 1e-06,
|
||||
|
|
@ -36652,7 +36664,9 @@
|
|||
"audio"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_audio_output": true
|
||||
"supports_audio_output": true,
|
||||
"tpm": 250000,
|
||||
"rpm": 10
|
||||
},
|
||||
"gemini/gemini-2.5-flash-native-audio-preview-12-2025": {
|
||||
"input_cost_per_audio_token": 1e-06,
|
||||
|
|
@ -36676,7 +36690,9 @@
|
|||
"audio"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_audio_output": true
|
||||
"supports_audio_output": true,
|
||||
"tpm": 250000,
|
||||
"rpm": 10
|
||||
},
|
||||
"gemini-2.5-flash-preview-tts": {
|
||||
"input_cost_per_token": 3e-07,
|
||||
|
|
|
|||
22
poetry.lock
generated
22
poetry.lock
generated
|
|
@ -1,4 +1,4 @@
|
|||
# This file is automatically @generated by Poetry 2.1.4 and should not be changed by hand.
|
||||
# This file is automatically @generated by Poetry 2.2.1 and should not be changed by hand.
|
||||
|
||||
[[package]]
|
||||
name = "a2a-sdk"
|
||||
|
|
@ -5659,6 +5659,24 @@ files = [
|
|||
[package.extras]
|
||||
dev = ["build", "flake8", "mypy", "pytest", "twine"]
|
||||
|
||||
[[package]]
|
||||
name = "pyroscope-io"
|
||||
version = "0.8.16"
|
||||
description = "Pyroscope Python integration"
|
||||
optional = false
|
||||
python-versions = "*"
|
||||
groups = ["main"]
|
||||
markers = "extra == \"proxy\" and sys_platform != \"win32\""
|
||||
files = [
|
||||
{file = "pyroscope_io-0.8.16-py2.py3-none-macosx_11_0_arm64.whl", hash = "sha256:e07edcfd59f5bdce42948b92c9b118c824edbd551730305f095a6b9af401a9e8"},
|
||||
{file = "pyroscope_io-0.8.16-py2.py3-none-macosx_11_0_x86_64.whl", hash = "sha256:dc98355e27c0b7b61f27066500fe1045b70e9459bb8b9a3082bc4755cb6392b6"},
|
||||
{file = "pyroscope_io-0.8.16-py2.py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:86f0f047554ff62bd92c3e5a26bc2809ccd467d11fbacb9fef898ba299dbda59"},
|
||||
{file = "pyroscope_io-0.8.16-py2.py3-none-manylinux2014_x86_64.manylinux_2_17_x86_64.whl", hash = "sha256:6b91ce5b240f8de756c16a17022ca8e25ef8a4eed461c7d074b8a0841cf7b445"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
cffi = ">=1.6.0"
|
||||
|
||||
[[package]]
|
||||
name = "pytest"
|
||||
version = "7.4.4"
|
||||
|
|
@ -8516,7 +8534,7 @@ extra-proxy = ["a2a-sdk", "azure-identity", "azure-keyvault-secrets", "google-cl
|
|||
google = ["google-cloud-aiplatform"]
|
||||
grpc = ["grpcio", "grpcio"]
|
||||
mlflow = ["mlflow"]
|
||||
proxy = ["PyJWT", "apscheduler", "azure-identity", "azure-storage-blob", "backoff", "boto3", "cryptography", "fastapi", "fastapi-sso", "gunicorn", "litellm-enterprise", "litellm-proxy-extras", "mcp", "orjson", "polars", "pynacl", "python-multipart", "pyyaml", "rich", "rq", "soundfile", "uvicorn", "uvloop", "websockets"]
|
||||
proxy = ["PyJWT", "apscheduler", "azure-identity", "azure-storage-blob", "backoff", "boto3", "cryptography", "fastapi", "fastapi-sso", "gunicorn", "litellm-enterprise", "litellm-proxy-extras", "mcp", "orjson", "polars", "pynacl", "pyroscope-io", "python-multipart", "pyyaml", "rich", "rq", "soundfile", "uvicorn", "uvloop", "websockets"]
|
||||
semantic-router = ["semantic-router"]
|
||||
utils = ["numpydoc"]
|
||||
|
||||
|
|
|
|||
|
|
@ -61,7 +61,7 @@ boto3 = { version = "1.40.76", optional = true }
|
|||
redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.9' and python_version < '3.14'"}
|
||||
mcp = {version = ">=1.25.0,<2.0.0", optional = true, python = ">=3.10"}
|
||||
a2a-sdk = {version = "^0.3.22", optional = true, python = ">=3.10"}
|
||||
litellm-proxy-extras = {version = "0.4.36", optional = true}
|
||||
litellm-proxy-extras = {version = "0.4.37", optional = true}
|
||||
rich = {version = "13.7.1", optional = true}
|
||||
litellm-enterprise = {version = "0.1.31", optional = true}
|
||||
diskcache = {version = "^5.6.1", optional = true}
|
||||
|
|
@ -69,6 +69,7 @@ polars = {version = "^1.31.0", optional = true, python = ">=3.10"}
|
|||
semantic-router = {version = ">=0.1.12", optional = true, python = ">=3.9,<3.14"}
|
||||
mlflow = {version = ">3.1.4", optional = true, python = ">=3.10"}
|
||||
soundfile = {version = "^0.12.1", optional = true}
|
||||
pyroscope-io = {version = "^0.8", optional = true, markers = "sys_platform != 'win32'"}
|
||||
# grpcio constraints:
|
||||
# - 1.62.3+ required by grpcio-status
|
||||
# - 1.68.0-1.68.1 has reconnect bug (https://github.com/grpc/grpc/issues/38290)
|
||||
|
|
@ -104,6 +105,7 @@ proxy = [
|
|||
"rich",
|
||||
"polars",
|
||||
"soundfile",
|
||||
"pyroscope-io",
|
||||
]
|
||||
|
||||
extra_proxy = [
|
||||
|
|
@ -121,6 +123,8 @@ utils = [
|
|||
"numpydoc",
|
||||
]
|
||||
|
||||
|
||||
|
||||
caching = ["diskcache"]
|
||||
|
||||
semantic-router = ["semantic-router"]
|
||||
|
|
|
|||
|
|
@ -55,7 +55,7 @@ grpcio>=1.75.0; python_version >= "3.14"
|
|||
sentry_sdk==2.21.0 # for sentry error handling
|
||||
detect-secrets==1.5.0 # Enterprise - secret detection / masking in LLM requests
|
||||
tzdata==2025.1 # IANA time zone database
|
||||
litellm-proxy-extras==0.4.36 # for proxy extras - e.g. prisma migrations
|
||||
litellm-proxy-extras==0.4.37 # for proxy extras - e.g. prisma migrations
|
||||
llm-sandbox==0.3.31 # for skill execution in sandbox
|
||||
### LITELLM PACKAGE DEPENDENCIES
|
||||
python-dotenv==1.0.1 # for env
|
||||
|
|
|
|||
|
|
@ -930,7 +930,7 @@ model LiteLLM_AccessGroupTable {
|
|||
description String?
|
||||
|
||||
// Resource memberships - explicit arrays per type
|
||||
access_model_ids String[] @default([])
|
||||
access_model_names String[] @default([])
|
||||
access_mcp_server_ids String[] @default([])
|
||||
access_agent_ids String[] @default([])
|
||||
|
||||
|
|
|
|||
|
|
@ -1278,3 +1278,86 @@ def test_transform_response_preserves_annotations():
|
|||
assert result.usage.total_tokens == 30
|
||||
|
||||
print("✓ Annotations from Responses API are correctly preserved in Chat Completions format")
|
||||
|
||||
|
||||
def test_convert_chat_completion_messages_to_responses_api_system_string():
|
||||
"""Test that string system content is extracted into instructions."""
|
||||
from litellm.completion_extras.litellm_responses_transformation.transformation import (
|
||||
LiteLLMResponsesTransformationHandler,
|
||||
)
|
||||
|
||||
handler = LiteLLMResponsesTransformationHandler()
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": "Hello"},
|
||||
]
|
||||
|
||||
input_items, instructions = handler.convert_chat_completion_messages_to_responses_api(messages)
|
||||
|
||||
assert instructions == "You are a helpful assistant."
|
||||
# System message should NOT appear in input items
|
||||
for item in input_items:
|
||||
assert item.get("role") != "system"
|
||||
# User message should be in input items
|
||||
assert len(input_items) == 1
|
||||
assert input_items[0]["role"] == "user"
|
||||
|
||||
|
||||
def test_convert_chat_completion_messages_to_responses_api_system_list_content():
|
||||
"""Test that list-format system content blocks are extracted into instructions.
|
||||
|
||||
This happens when requests arrive via the Anthropic /v1/messages adapter,
|
||||
which converts system prompts into list-format content blocks.
|
||||
"""
|
||||
from litellm.completion_extras.litellm_responses_transformation.transformation import (
|
||||
LiteLLMResponsesTransformationHandler,
|
||||
)
|
||||
|
||||
handler = LiteLLMResponsesTransformationHandler()
|
||||
|
||||
messages = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{"type": "text", "text": "You are a helpful assistant."},
|
||||
{"type": "text", "text": "Be concise."},
|
||||
],
|
||||
},
|
||||
{"role": "user", "content": "Hello"},
|
||||
]
|
||||
|
||||
input_items, instructions = handler.convert_chat_completion_messages_to_responses_api(messages)
|
||||
|
||||
assert instructions == "You are a helpful assistant. Be concise."
|
||||
# System message should NOT appear in input items
|
||||
for item in input_items:
|
||||
assert item.get("role") != "system"
|
||||
assert len(input_items) == 1
|
||||
assert input_items[0]["role"] == "user"
|
||||
|
||||
|
||||
def test_convert_chat_completion_messages_to_responses_api_multiple_system_messages():
|
||||
"""Test that multiple system messages (string and list) are concatenated."""
|
||||
from litellm.completion_extras.litellm_responses_transformation.transformation import (
|
||||
LiteLLMResponsesTransformationHandler,
|
||||
)
|
||||
|
||||
handler = LiteLLMResponsesTransformationHandler()
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{
|
||||
"role": "system",
|
||||
"content": [
|
||||
{"type": "text", "text": "Be concise."},
|
||||
],
|
||||
},
|
||||
{"role": "user", "content": "Hello"},
|
||||
]
|
||||
|
||||
input_items, instructions = handler.convert_chat_completion_messages_to_responses_api(messages)
|
||||
|
||||
assert instructions == "You are a helpful assistant. Be concise."
|
||||
for item in input_items:
|
||||
assert item.get("role") != "system"
|
||||
|
|
|
|||
|
|
@ -422,3 +422,77 @@ async def test_return_user_api_key_auth_obj_user_spend_and_budget():
|
|||
assert result.user_tpm_limit == 1000
|
||||
assert result.user_rpm_limit == 100
|
||||
assert result.user_email == "test@example.com"
|
||||
|
||||
|
||||
def test_proxy_admin_jwt_auth_includes_identity_fields():
|
||||
"""
|
||||
Test that the proxy admin early-return path in JWT auth populates
|
||||
user_id, team_id, team_alias, team_metadata, org_id, and end_user_id.
|
||||
|
||||
Regression test: previously the is_proxy_admin branch only set user_role
|
||||
and parent_otel_span, discarding all identity fields resolved from the JWT.
|
||||
This caused blank Team Name and Internal User in Request Logs UI.
|
||||
"""
|
||||
from litellm.proxy._types import LiteLLM_TeamTable, LitellmUserRoles, UserAPIKeyAuth
|
||||
|
||||
team_object = LiteLLM_TeamTable(
|
||||
team_id="team-123",
|
||||
team_alias="my-team",
|
||||
metadata={"tags": ["prod"], "env": "production"},
|
||||
)
|
||||
|
||||
# Simulate the proxy admin early-return path (user_api_key_auth.py ~line 586)
|
||||
result = UserAPIKeyAuth(
|
||||
api_key=None,
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
user_id="user-abc",
|
||||
team_id="team-123",
|
||||
team_alias=(
|
||||
team_object.team_alias if team_object is not None else None
|
||||
),
|
||||
team_metadata=team_object.metadata if team_object is not None else None,
|
||||
org_id="org-456",
|
||||
end_user_id="end-user-789",
|
||||
parent_otel_span=None,
|
||||
)
|
||||
|
||||
assert result.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
assert result.user_id == "user-abc"
|
||||
assert result.team_id == "team-123"
|
||||
assert result.team_alias == "my-team"
|
||||
assert result.team_metadata == {"tags": ["prod"], "env": "production"}
|
||||
assert result.org_id == "org-456"
|
||||
assert result.end_user_id == "end-user-789"
|
||||
assert result.api_key is None
|
||||
|
||||
|
||||
def test_proxy_admin_jwt_auth_handles_no_team_object():
|
||||
"""
|
||||
Test that the proxy admin early-return path works correctly when
|
||||
team_object is None (user has admin role but no team association).
|
||||
"""
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
|
||||
team_object = None
|
||||
|
||||
result = UserAPIKeyAuth(
|
||||
api_key=None,
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
user_id="admin-user",
|
||||
team_id=None,
|
||||
team_alias=(
|
||||
team_object.team_alias if team_object is not None else None
|
||||
),
|
||||
team_metadata=team_object.metadata if team_object is not None else None,
|
||||
org_id=None,
|
||||
end_user_id=None,
|
||||
parent_otel_span=None,
|
||||
)
|
||||
|
||||
assert result.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
assert result.user_id == "admin-user"
|
||||
assert result.team_id is None
|
||||
assert result.team_alias is None
|
||||
assert result.team_metadata is None
|
||||
assert result.org_id is None
|
||||
assert result.end_user_id is None
|
||||
|
|
|
|||
|
|
@ -28,7 +28,7 @@ def _make_access_group_record(
|
|||
access_group_id: str = "ag-123",
|
||||
access_group_name: str = "test-group",
|
||||
description: str | None = "Test description",
|
||||
access_model_ids: list | None = None,
|
||||
access_model_names: list | None = None,
|
||||
access_mcp_server_ids: list | None = None,
|
||||
access_agent_ids: list | None = None,
|
||||
assigned_team_ids: list | None = None,
|
||||
|
|
@ -41,7 +41,7 @@ def _make_access_group_record(
|
|||
record.access_group_id = access_group_id
|
||||
record.access_group_name = access_group_name
|
||||
record.description = description
|
||||
record.access_model_ids = access_model_ids or []
|
||||
record.access_model_names = access_model_names or []
|
||||
record.access_mcp_server_ids = access_mcp_server_ids or []
|
||||
record.access_agent_ids = access_agent_ids or []
|
||||
record.assigned_team_ids = assigned_team_ids or []
|
||||
|
|
@ -64,7 +64,7 @@ def client_and_mocks(monkeypatch):
|
|||
access_group_id="ag-new",
|
||||
access_group_name=data.get("access_group_name", "new"),
|
||||
description=data.get("description"),
|
||||
access_model_ids=data.get("access_model_ids", []),
|
||||
access_model_names=data.get("access_model_names", []),
|
||||
access_mcp_server_ids=data.get("access_mcp_server_ids", []),
|
||||
access_agent_ids=data.get("access_agent_ids", []),
|
||||
assigned_team_ids=data.get("assigned_team_ids", []),
|
||||
|
|
@ -80,7 +80,7 @@ def client_and_mocks(monkeypatch):
|
|||
access_group_id=where.get("access_group_id", "ag-123"),
|
||||
access_group_name=data.get("access_group_name", "updated"),
|
||||
description=data.get("description"),
|
||||
access_model_ids=data.get("access_model_ids", []),
|
||||
access_model_names=data.get("access_model_names", []),
|
||||
access_mcp_server_ids=data.get("access_mcp_server_ids", []),
|
||||
access_agent_ids=data.get("access_agent_ids", []),
|
||||
assigned_team_ids=data.get("assigned_team_ids", []),
|
||||
|
|
@ -147,7 +147,7 @@ ACCESS_GROUP_PATHS = ["/v1/access_group", "/v1/unified_access_group"]
|
|||
{
|
||||
"access_group_name": "group-b",
|
||||
"description": "Group B description",
|
||||
"access_model_ids": ["model-1"],
|
||||
"access_model_names": ["model-1"],
|
||||
"access_mcp_server_ids": ["mcp-1"],
|
||||
"assigned_team_ids": ["team-1"],
|
||||
},
|
||||
|
|
@ -369,7 +369,7 @@ def test_get_access_group_forbidden_non_admin(client_and_mocks, user_role):
|
|||
"update_payload",
|
||||
[
|
||||
{"description": "Updated description"},
|
||||
{"access_model_ids": ["model-1", "model-2"]},
|
||||
{"access_model_names": ["model-1", "model-2"]},
|
||||
{"assigned_team_ids": [], "assigned_key_ids": ["key-1"]},
|
||||
],
|
||||
)
|
||||
|
|
@ -431,6 +431,57 @@ def test_update_access_group_empty_body(client_and_mocks):
|
|||
assert call_kwargs["data"]["updated_by"] == "admin_user"
|
||||
|
||||
|
||||
def test_update_access_group_name_success(client_and_mocks):
|
||||
"""Update access_group_name succeeds when new name is unique."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
|
||||
existing = _make_access_group_record(access_group_id="ag-update", access_group_name="old-name")
|
||||
mock_table.find_unique = AsyncMock(return_value=existing)
|
||||
|
||||
resp = client.put("/v1/access_group/ag-update", json={"access_group_name": "new-name"})
|
||||
assert resp.status_code == 200
|
||||
mock_table.update.assert_awaited_once()
|
||||
call_kwargs = mock_table.update.call_args.kwargs
|
||||
assert call_kwargs["data"]["access_group_name"] == "new-name"
|
||||
|
||||
|
||||
def test_update_access_group_name_duplicate_conflict(client_and_mocks):
|
||||
"""Update access_group_name to existing name returns 409 (unique constraint)."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
|
||||
existing = _make_access_group_record(access_group_id="ag-update", access_group_name="old-name")
|
||||
mock_table.find_unique = AsyncMock(return_value=existing)
|
||||
mock_table.update = AsyncMock(
|
||||
side_effect=Exception("Unique constraint failed on the fields: (`access_group_name`)")
|
||||
)
|
||||
|
||||
resp = client.put("/v1/access_group/ag-update", json={"access_group_name": "taken-name"})
|
||||
assert resp.status_code == 409
|
||||
assert "already exists" in resp.json()["detail"]
|
||||
mock_table.update.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"error_message",
|
||||
[
|
||||
"Unique constraint failed on the fields: (`access_group_name`)",
|
||||
"P2002: Unique constraint failed",
|
||||
"unique constraint violation",
|
||||
],
|
||||
)
|
||||
def test_update_access_group_name_unique_constraint_returns_409(client_and_mocks, error_message):
|
||||
"""Update access_group_name: Prisma unique constraint surfaces as 409."""
|
||||
client, _, mock_table = client_and_mocks
|
||||
|
||||
existing = _make_access_group_record(access_group_id="ag-update", access_group_name="old-name")
|
||||
mock_table.find_unique = AsyncMock(return_value=existing)
|
||||
mock_table.update = AsyncMock(side_effect=Exception(error_message))
|
||||
|
||||
resp = client.put("/v1/access_group/ag-update", json={"access_group_name": "race-name"})
|
||||
assert resp.status_code == 409
|
||||
assert "already exists" in resp.json()["detail"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# DELETE
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -2087,6 +2087,143 @@ async def test_add_litellm_data_to_request_adds_headers_to_metadata():
|
|||
assert "headers" in result["proxy_server_request"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_pass_through_route_custom_body_url_target():
|
||||
"""
|
||||
Test that the URL-based endpoint_func created by create_pass_through_route
|
||||
accepts a custom_body parameter and forwards it to pass_through_request,
|
||||
taking precedence over the request-parsed body.
|
||||
|
||||
This verifies the fix for issue #16999 where bedrock_proxy_route passes
|
||||
custom_body=data to the endpoint function, which previously crashed with:
|
||||
TypeError: endpoint_func() got an unexpected keyword argument 'custom_body'
|
||||
"""
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
||||
create_pass_through_route,
|
||||
)
|
||||
|
||||
unique_path = "/test/path/unique/custom_body_url"
|
||||
endpoint_func = create_pass_through_route(
|
||||
endpoint=unique_path,
|
||||
target="https://bedrock-agent-runtime.us-east-1.amazonaws.com",
|
||||
custom_headers={"Content-Type": "application/json"},
|
||||
_forward_headers=True,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_request"
|
||||
) as mock_pass_through, patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.InitPassThroughEndpointHelpers.is_registered_pass_through_route"
|
||||
) as mock_is_registered, patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.InitPassThroughEndpointHelpers.get_registered_pass_through_route"
|
||||
) as mock_get_registered, patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._parse_request_data_by_content_type"
|
||||
) as mock_parse_request:
|
||||
mock_pass_through.return_value = MagicMock()
|
||||
mock_is_registered.return_value = True
|
||||
mock_get_registered.return_value = None
|
||||
# Simulate the request parser returning a different body
|
||||
mock_parse_request.return_value = (
|
||||
{}, # query_params_data
|
||||
{"parsed_from_request": True}, # custom_body_data (from request)
|
||||
None, # file_data
|
||||
False, # stream
|
||||
)
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.url = MagicMock()
|
||||
mock_request.url.path = unique_path
|
||||
mock_request.path_params = {}
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
||||
mock_user_api_key_dict = MagicMock()
|
||||
mock_user_api_key_dict.api_key = "test-key"
|
||||
|
||||
# The caller-supplied body (e.g. from bedrock_proxy_route)
|
||||
bedrock_body = {
|
||||
"retrievalQuery": {"text": "What is in the knowledge base?"},
|
||||
}
|
||||
|
||||
# Call endpoint_func with custom_body — this is the call that
|
||||
# used to crash with TypeError before the fix
|
||||
await endpoint_func(
|
||||
request=mock_request,
|
||||
fastapi_response=MagicMock(),
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
custom_body=bedrock_body,
|
||||
)
|
||||
|
||||
mock_pass_through.assert_called_once()
|
||||
call_kwargs = mock_pass_through.call_args[1]
|
||||
|
||||
# The critical assertion: custom_body takes precedence over
|
||||
# the body parsed from the raw request
|
||||
assert call_kwargs["custom_body"] == bedrock_body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_pass_through_route_no_custom_body_falls_back():
|
||||
"""
|
||||
Test that the URL-based endpoint_func falls back to the request-parsed body
|
||||
when custom_body is not provided.
|
||||
|
||||
This ensures the default pass-through behavior is preserved — only the
|
||||
Bedrock proxy route (and similar callers) supply a pre-built body.
|
||||
"""
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
||||
create_pass_through_route,
|
||||
)
|
||||
|
||||
unique_path = "/test/path/unique/no_custom_body"
|
||||
endpoint_func = create_pass_through_route(
|
||||
endpoint=unique_path,
|
||||
target="http://example.com/api",
|
||||
custom_headers={},
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_request"
|
||||
) as mock_pass_through, patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.InitPassThroughEndpointHelpers.is_registered_pass_through_route"
|
||||
) as mock_is_registered, patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.InitPassThroughEndpointHelpers.get_registered_pass_through_route"
|
||||
) as mock_get_registered, patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._parse_request_data_by_content_type"
|
||||
) as mock_parse_request:
|
||||
mock_pass_through.return_value = MagicMock()
|
||||
mock_is_registered.return_value = True
|
||||
mock_get_registered.return_value = None
|
||||
request_parsed_body = {"key": "from_request"}
|
||||
mock_parse_request.return_value = (
|
||||
{}, # query_params_data
|
||||
request_parsed_body, # custom_body_data
|
||||
None, # file_data
|
||||
False, # stream
|
||||
)
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.url = MagicMock()
|
||||
mock_request.url.path = unique_path
|
||||
mock_request.path_params = {}
|
||||
mock_request.query_params = QueryParams({})
|
||||
|
||||
mock_user_api_key_dict = MagicMock()
|
||||
mock_user_api_key_dict.api_key = "test-key"
|
||||
|
||||
# Call without custom_body — should use the request-parsed body
|
||||
await endpoint_func(
|
||||
request=mock_request,
|
||||
fastapi_response=MagicMock(),
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
|
||||
mock_pass_through.assert_called_once()
|
||||
call_kwargs = mock_pass_through.call_args[1]
|
||||
|
||||
# Should fall back to the body parsed from the request
|
||||
assert call_kwargs["custom_body"] == request_parsed_body
|
||||
|
||||
|
||||
def test_build_full_path_with_root_default():
|
||||
"""
|
||||
Test _build_full_path_with_root with default root path (/)
|
||||
|
|
|
|||
484
tests/test_litellm/proxy/policy_engine/test_pipeline_executor.py
Normal file
484
tests/test_litellm/proxy/policy_engine/test_pipeline_executor.py
Normal file
|
|
@ -0,0 +1,484 @@
|
|||
"""
|
||||
Tests for the pipeline executor.
|
||||
|
||||
Uses mock guardrails to validate pipeline execution without external services.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor
|
||||
from litellm.types.proxy.policy_engine.pipeline_types import (
|
||||
GuardrailPipeline,
|
||||
PipelineStep,
|
||||
)
|
||||
|
||||
try:
|
||||
from fastapi.exceptions import HTTPException
|
||||
except ImportError:
|
||||
HTTPException = None
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# Mock Guardrails
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class AlwaysFailGuardrail(CustomGuardrail):
|
||||
"""Mock guardrail that always raises HTTPException(400)."""
|
||||
|
||||
def __init__(self, guardrail_name: str):
|
||||
super().__init__(
|
||||
guardrail_name=guardrail_name,
|
||||
event_hook="pre_call",
|
||||
default_on=True,
|
||||
)
|
||||
self.calls = 0
|
||||
|
||||
def should_run_guardrail(self, data, event_type) -> bool:
|
||||
return True
|
||||
|
||||
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
|
||||
self.calls += 1
|
||||
raise HTTPException(status_code=400, detail="Content policy violation")
|
||||
|
||||
|
||||
class AlwaysPassGuardrail(CustomGuardrail):
|
||||
"""Mock guardrail that always passes."""
|
||||
|
||||
def __init__(self, guardrail_name: str):
|
||||
super().__init__(
|
||||
guardrail_name=guardrail_name,
|
||||
event_hook="pre_call",
|
||||
default_on=True,
|
||||
)
|
||||
self.calls = 0
|
||||
|
||||
def should_run_guardrail(self, data, event_type) -> bool:
|
||||
return True
|
||||
|
||||
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
|
||||
self.calls += 1
|
||||
return None
|
||||
|
||||
|
||||
class PiiMaskingGuardrail(CustomGuardrail):
|
||||
"""Mock guardrail that masks PII in messages and returns modified data."""
|
||||
|
||||
def __init__(self, guardrail_name: str):
|
||||
super().__init__(
|
||||
guardrail_name=guardrail_name,
|
||||
event_hook="pre_call",
|
||||
default_on=True,
|
||||
)
|
||||
self.calls = 0
|
||||
self.received_messages = None
|
||||
|
||||
def should_run_guardrail(self, data, event_type) -> bool:
|
||||
return True
|
||||
|
||||
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
|
||||
self.calls += 1
|
||||
self.received_messages = data.get("messages", [])
|
||||
masked_messages = []
|
||||
for msg in data.get("messages", []):
|
||||
masked_msg = dict(msg)
|
||||
masked_msg["content"] = msg["content"].replace(
|
||||
"John Smith", "[REDACTED]"
|
||||
)
|
||||
masked_messages.append(masked_msg)
|
||||
return {"messages": masked_messages}
|
||||
|
||||
|
||||
class ContentCheckGuardrail(CustomGuardrail):
|
||||
"""Mock guardrail that records what messages it received."""
|
||||
|
||||
def __init__(self, guardrail_name: str):
|
||||
super().__init__(
|
||||
guardrail_name=guardrail_name,
|
||||
event_hook="pre_call",
|
||||
default_on=True,
|
||||
)
|
||||
self.calls = 0
|
||||
self.received_messages = None
|
||||
|
||||
def should_run_guardrail(self, data, event_type) -> bool:
|
||||
return True
|
||||
|
||||
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
|
||||
self.calls += 1
|
||||
self.received_messages = data.get("messages", [])
|
||||
return None
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
# Tests
|
||||
# ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.skipif(HTTPException is None, reason="fastapi not installed")
|
||||
@pytest.mark.asyncio
|
||||
async def test_escalation_step1_fails_step2_blocks():
|
||||
"""
|
||||
Pipeline: simple-filter (on_fail: next) -> advanced-filter (on_fail: block)
|
||||
Input: request that fails simple-filter
|
||||
Expected: simple-filter fails -> escalate -> advanced-filter fails -> block
|
||||
"""
|
||||
simple_guard = AlwaysFailGuardrail(guardrail_name="simple-filter")
|
||||
advanced_guard = AlwaysFailGuardrail(guardrail_name="advanced-filter")
|
||||
|
||||
pipeline = GuardrailPipeline(
|
||||
mode="pre_call",
|
||||
steps=[
|
||||
PipelineStep(
|
||||
guardrail="simple-filter", on_fail="next", on_pass="allow"
|
||||
),
|
||||
PipelineStep(
|
||||
guardrail="advanced-filter", on_fail="block", on_pass="allow"
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
original_callbacks = litellm.callbacks.copy()
|
||||
litellm.callbacks = [simple_guard, advanced_guard]
|
||||
|
||||
try:
|
||||
result = await PipelineExecutor.execute_steps(
|
||||
steps=pipeline.steps,
|
||||
mode=pipeline.mode,
|
||||
data={"messages": [{"role": "user", "content": "bad content"}]},
|
||||
user_api_key_dict=MagicMock(),
|
||||
call_type="completion",
|
||||
policy_name="content-safety",
|
||||
)
|
||||
|
||||
assert simple_guard.calls == 1
|
||||
assert advanced_guard.calls == 1
|
||||
assert result.terminal_action == "block"
|
||||
assert len(result.step_results) == 2
|
||||
assert result.step_results[0].guardrail_name == "simple-filter"
|
||||
assert result.step_results[0].outcome == "fail"
|
||||
assert result.step_results[0].action_taken == "next"
|
||||
assert result.step_results[1].guardrail_name == "advanced-filter"
|
||||
assert result.step_results[1].outcome == "fail"
|
||||
assert result.step_results[1].action_taken == "block"
|
||||
finally:
|
||||
litellm.callbacks = original_callbacks
|
||||
|
||||
|
||||
@pytest.mark.skipif(HTTPException is None, reason="fastapi not installed")
|
||||
@pytest.mark.asyncio
|
||||
async def test_early_allow_step1_passes_step2_skipped():
|
||||
"""
|
||||
Pipeline: simple-filter (on_pass: allow) -> advanced-filter
|
||||
Input: clean request that passes simple-filter
|
||||
Expected: simple-filter passes -> allow (advanced-filter never called)
|
||||
"""
|
||||
simple_guard = AlwaysPassGuardrail(guardrail_name="simple-filter")
|
||||
advanced_guard = AlwaysFailGuardrail(guardrail_name="advanced-filter")
|
||||
|
||||
pipeline = GuardrailPipeline(
|
||||
mode="pre_call",
|
||||
steps=[
|
||||
PipelineStep(
|
||||
guardrail="simple-filter", on_fail="next", on_pass="allow"
|
||||
),
|
||||
PipelineStep(
|
||||
guardrail="advanced-filter", on_fail="block", on_pass="allow"
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
original_callbacks = litellm.callbacks.copy()
|
||||
litellm.callbacks = [simple_guard, advanced_guard]
|
||||
|
||||
try:
|
||||
result = await PipelineExecutor.execute_steps(
|
||||
steps=pipeline.steps,
|
||||
mode=pipeline.mode,
|
||||
data={"messages": [{"role": "user", "content": "clean content"}]},
|
||||
user_api_key_dict=MagicMock(),
|
||||
call_type="completion",
|
||||
policy_name="content-safety",
|
||||
)
|
||||
|
||||
assert simple_guard.calls == 1
|
||||
assert advanced_guard.calls == 0
|
||||
assert result.terminal_action == "allow"
|
||||
assert len(result.step_results) == 1
|
||||
assert result.step_results[0].outcome == "pass"
|
||||
assert result.step_results[0].action_taken == "allow"
|
||||
finally:
|
||||
litellm.callbacks = original_callbacks
|
||||
|
||||
|
||||
@pytest.mark.skipif(HTTPException is None, reason="fastapi not installed")
|
||||
@pytest.mark.asyncio
|
||||
async def test_escalation_step1_fails_step2_passes():
|
||||
"""
|
||||
Pipeline: simple-filter (on_fail: next) -> advanced-filter (on_pass: allow)
|
||||
Input: request that fails simple but passes advanced
|
||||
Expected: simple-filter fails -> escalate -> advanced-filter passes -> allow
|
||||
"""
|
||||
simple_guard = AlwaysFailGuardrail(guardrail_name="simple-filter")
|
||||
advanced_guard = AlwaysPassGuardrail(guardrail_name="advanced-filter")
|
||||
|
||||
pipeline = GuardrailPipeline(
|
||||
mode="pre_call",
|
||||
steps=[
|
||||
PipelineStep(
|
||||
guardrail="simple-filter", on_fail="next", on_pass="allow"
|
||||
),
|
||||
PipelineStep(
|
||||
guardrail="advanced-filter", on_fail="block", on_pass="allow"
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
original_callbacks = litellm.callbacks.copy()
|
||||
litellm.callbacks = [simple_guard, advanced_guard]
|
||||
|
||||
try:
|
||||
result = await PipelineExecutor.execute_steps(
|
||||
steps=pipeline.steps,
|
||||
mode=pipeline.mode,
|
||||
data={"messages": [{"role": "user", "content": "borderline content"}]},
|
||||
user_api_key_dict=MagicMock(),
|
||||
call_type="completion",
|
||||
policy_name="content-safety",
|
||||
)
|
||||
|
||||
assert simple_guard.calls == 1
|
||||
assert advanced_guard.calls == 1
|
||||
assert result.terminal_action == "allow"
|
||||
assert len(result.step_results) == 2
|
||||
assert result.step_results[0].outcome == "fail"
|
||||
assert result.step_results[0].action_taken == "next"
|
||||
assert result.step_results[1].outcome == "pass"
|
||||
assert result.step_results[1].action_taken == "allow"
|
||||
finally:
|
||||
litellm.callbacks = original_callbacks
|
||||
|
||||
|
||||
@pytest.mark.skipif(HTTPException is None, reason="fastapi not installed")
|
||||
@pytest.mark.asyncio
|
||||
async def test_data_forwarding_pii_masking():
|
||||
"""
|
||||
Pipeline: pii-masker (pass_data: true, on_pass: next) -> content-check (on_pass: allow)
|
||||
Input: "Hello John Smith"
|
||||
Expected: pii-masker masks -> content-check receives "[REDACTED]" -> allow
|
||||
"""
|
||||
pii_guard = PiiMaskingGuardrail(guardrail_name="pii-masker")
|
||||
content_guard = ContentCheckGuardrail(guardrail_name="content-check")
|
||||
|
||||
pipeline = GuardrailPipeline(
|
||||
mode="pre_call",
|
||||
steps=[
|
||||
PipelineStep(
|
||||
guardrail="pii-masker",
|
||||
on_fail="block",
|
||||
on_pass="next",
|
||||
pass_data=True,
|
||||
),
|
||||
PipelineStep(
|
||||
guardrail="content-check", on_fail="block", on_pass="allow"
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
original_callbacks = litellm.callbacks.copy()
|
||||
litellm.callbacks = [pii_guard, content_guard]
|
||||
|
||||
try:
|
||||
result = await PipelineExecutor.execute_steps(
|
||||
steps=pipeline.steps,
|
||||
mode=pipeline.mode,
|
||||
data={
|
||||
"messages": [{"role": "user", "content": "Hello John Smith"}]
|
||||
},
|
||||
user_api_key_dict=MagicMock(),
|
||||
call_type="completion",
|
||||
policy_name="pii-then-safety",
|
||||
)
|
||||
|
||||
assert pii_guard.calls == 1
|
||||
assert content_guard.calls == 1
|
||||
assert content_guard.received_messages[0]["content"] == "Hello [REDACTED]"
|
||||
assert result.terminal_action == "allow"
|
||||
assert result.modified_data is not None
|
||||
assert result.modified_data["messages"][0]["content"] == "Hello [REDACTED]"
|
||||
finally:
|
||||
litellm.callbacks = original_callbacks
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_guardrail_not_found_uses_on_fail():
|
||||
"""
|
||||
If a guardrail is not found, treat as error and use on_fail action.
|
||||
"""
|
||||
pipeline = GuardrailPipeline(
|
||||
mode="pre_call",
|
||||
steps=[
|
||||
PipelineStep(
|
||||
guardrail="nonexistent-guard",
|
||||
on_fail="block",
|
||||
on_pass="allow",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
original_callbacks = litellm.callbacks.copy()
|
||||
litellm.callbacks = []
|
||||
|
||||
try:
|
||||
result = await PipelineExecutor.execute_steps(
|
||||
steps=pipeline.steps,
|
||||
mode=pipeline.mode,
|
||||
data={"messages": [{"role": "user", "content": "test"}]},
|
||||
user_api_key_dict=MagicMock(),
|
||||
call_type="completion",
|
||||
policy_name="test-policy",
|
||||
)
|
||||
|
||||
assert result.terminal_action == "block"
|
||||
assert result.step_results[0].outcome == "error"
|
||||
assert "not found" in result.step_results[0].error_detail
|
||||
finally:
|
||||
litellm.callbacks = original_callbacks
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_guardrail_not_found_with_next_continues():
|
||||
"""
|
||||
If a guardrail is not found and on_fail is 'next', continue to next step.
|
||||
"""
|
||||
pass_guard = AlwaysPassGuardrail(guardrail_name="fallback-guard")
|
||||
|
||||
pipeline = GuardrailPipeline(
|
||||
mode="pre_call",
|
||||
steps=[
|
||||
PipelineStep(
|
||||
guardrail="nonexistent-guard",
|
||||
on_fail="next",
|
||||
on_pass="allow",
|
||||
),
|
||||
PipelineStep(
|
||||
guardrail="fallback-guard",
|
||||
on_fail="block",
|
||||
on_pass="allow",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
original_callbacks = litellm.callbacks.copy()
|
||||
litellm.callbacks = [pass_guard]
|
||||
|
||||
try:
|
||||
result = await PipelineExecutor.execute_steps(
|
||||
steps=pipeline.steps,
|
||||
mode=pipeline.mode,
|
||||
data={"messages": [{"role": "user", "content": "test"}]},
|
||||
user_api_key_dict=MagicMock(),
|
||||
call_type="completion",
|
||||
policy_name="test-policy",
|
||||
)
|
||||
|
||||
assert result.terminal_action == "allow"
|
||||
assert len(result.step_results) == 2
|
||||
assert result.step_results[0].outcome == "error"
|
||||
assert result.step_results[0].action_taken == "next"
|
||||
assert result.step_results[1].outcome == "pass"
|
||||
assert pass_guard.calls == 1
|
||||
finally:
|
||||
litellm.callbacks = original_callbacks
|
||||
|
||||
|
||||
@pytest.mark.skipif(HTTPException is None, reason="fastapi not installed")
|
||||
@pytest.mark.asyncio
|
||||
async def test_single_step_pipeline_block():
|
||||
"""Single step pipeline that blocks."""
|
||||
guard = AlwaysFailGuardrail(guardrail_name="blocker")
|
||||
|
||||
pipeline = GuardrailPipeline(
|
||||
mode="pre_call",
|
||||
steps=[PipelineStep(guardrail="blocker", on_fail="block")],
|
||||
)
|
||||
|
||||
original_callbacks = litellm.callbacks.copy()
|
||||
litellm.callbacks = [guard]
|
||||
|
||||
try:
|
||||
result = await PipelineExecutor.execute_steps(
|
||||
steps=pipeline.steps,
|
||||
mode=pipeline.mode,
|
||||
data={"messages": [{"role": "user", "content": "test"}]},
|
||||
user_api_key_dict=MagicMock(),
|
||||
call_type="completion",
|
||||
policy_name="test",
|
||||
)
|
||||
|
||||
assert result.terminal_action == "block"
|
||||
assert guard.calls == 1
|
||||
finally:
|
||||
litellm.callbacks = original_callbacks
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_single_step_pipeline_allow():
|
||||
"""Single step pipeline that allows."""
|
||||
guard = AlwaysPassGuardrail(guardrail_name="passer")
|
||||
|
||||
pipeline = GuardrailPipeline(
|
||||
mode="pre_call",
|
||||
steps=[PipelineStep(guardrail="passer", on_pass="allow")],
|
||||
)
|
||||
|
||||
original_callbacks = litellm.callbacks.copy()
|
||||
litellm.callbacks = [guard]
|
||||
|
||||
try:
|
||||
result = await PipelineExecutor.execute_steps(
|
||||
steps=pipeline.steps,
|
||||
mode=pipeline.mode,
|
||||
data={"messages": [{"role": "user", "content": "test"}]},
|
||||
user_api_key_dict=MagicMock(),
|
||||
call_type="completion",
|
||||
policy_name="test",
|
||||
)
|
||||
|
||||
assert result.terminal_action == "allow"
|
||||
assert guard.calls == 1
|
||||
finally:
|
||||
litellm.callbacks = original_callbacks
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_step_results_include_duration():
|
||||
"""Step results should include timing information."""
|
||||
guard = AlwaysPassGuardrail(guardrail_name="timed")
|
||||
|
||||
pipeline = GuardrailPipeline(
|
||||
mode="pre_call",
|
||||
steps=[PipelineStep(guardrail="timed")],
|
||||
)
|
||||
|
||||
original_callbacks = litellm.callbacks.copy()
|
||||
litellm.callbacks = [guard]
|
||||
|
||||
try:
|
||||
result = await PipelineExecutor.execute_steps(
|
||||
steps=pipeline.steps,
|
||||
mode=pipeline.mode,
|
||||
data={"messages": [{"role": "user", "content": "test"}]},
|
||||
user_api_key_dict=MagicMock(),
|
||||
call_type="completion",
|
||||
policy_name="test",
|
||||
)
|
||||
|
||||
assert result.step_results[0].duration_seconds is not None
|
||||
assert result.step_results[0].duration_seconds >= 0
|
||||
finally:
|
||||
litellm.callbacks = original_callbacks
|
||||
138
tests/test_litellm/proxy/test_pyroscope.py
Normal file
138
tests/test_litellm/proxy/test_pyroscope.py
Normal file
|
|
@ -0,0 +1,138 @@
|
|||
"""Unit tests for ProxyStartupEvent._init_pyroscope (Grafana Pyroscope profiling)."""
|
||||
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||||
|
||||
|
||||
def _mock_pyroscope_module():
|
||||
"""Return a mock module so 'import pyroscope' succeeds in _init_pyroscope."""
|
||||
m = MagicMock()
|
||||
m.configure = MagicMock()
|
||||
return m
|
||||
|
||||
|
||||
def test_init_pyroscope_returns_cleanly_when_disabled():
|
||||
"""When LITELLM_ENABLE_PYROSCOPE is false, _init_pyroscope returns without error."""
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.get_secret_bool",
|
||||
return_value=False,
|
||||
):
|
||||
ProxyStartupEvent._init_pyroscope()
|
||||
|
||||
|
||||
def test_init_pyroscope_raises_when_enabled_but_missing_app_name():
|
||||
"""When LITELLM_ENABLE_PYROSCOPE is true but PYROSCOPE_APP_NAME is not set, raises ValueError."""
|
||||
mock_pyroscope = _mock_pyroscope_module()
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.get_secret_bool",
|
||||
return_value=True,
|
||||
), patch.dict(
|
||||
sys.modules,
|
||||
{"pyroscope": mock_pyroscope},
|
||||
), patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
"PYROSCOPE_APP_NAME": "",
|
||||
"PYROSCOPE_SERVER_ADDRESS": "http://localhost:4040",
|
||||
},
|
||||
clear=False,
|
||||
):
|
||||
with pytest.raises(ValueError, match="PYROSCOPE_APP_NAME"):
|
||||
ProxyStartupEvent._init_pyroscope()
|
||||
|
||||
|
||||
def test_init_pyroscope_raises_when_enabled_but_missing_server_address():
|
||||
"""When LITELLM_ENABLE_PYROSCOPE is true but PYROSCOPE_SERVER_ADDRESS is not set, raises ValueError."""
|
||||
mock_pyroscope = _mock_pyroscope_module()
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.get_secret_bool",
|
||||
return_value=True,
|
||||
), patch.dict(
|
||||
sys.modules,
|
||||
{"pyroscope": mock_pyroscope},
|
||||
), patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
"PYROSCOPE_APP_NAME": "myapp",
|
||||
"PYROSCOPE_SERVER_ADDRESS": "",
|
||||
},
|
||||
clear=False,
|
||||
):
|
||||
with pytest.raises(ValueError, match="PYROSCOPE_SERVER_ADDRESS"):
|
||||
ProxyStartupEvent._init_pyroscope()
|
||||
|
||||
|
||||
def test_init_pyroscope_raises_when_sample_rate_invalid():
|
||||
"""When PYROSCOPE_SAMPLE_RATE is not a number, raises ValueError."""
|
||||
mock_pyroscope = _mock_pyroscope_module()
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.get_secret_bool",
|
||||
return_value=True,
|
||||
), patch.dict(
|
||||
sys.modules,
|
||||
{"pyroscope": mock_pyroscope},
|
||||
), patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
"PYROSCOPE_APP_NAME": "myapp",
|
||||
"PYROSCOPE_SERVER_ADDRESS": "http://localhost:4040",
|
||||
"PYROSCOPE_SAMPLE_RATE": "not-a-number",
|
||||
},
|
||||
clear=False,
|
||||
):
|
||||
with pytest.raises(ValueError, match="PYROSCOPE_SAMPLE_RATE"):
|
||||
ProxyStartupEvent._init_pyroscope()
|
||||
|
||||
|
||||
def test_init_pyroscope_accepts_integer_sample_rate():
|
||||
"""When enabled with valid config and integer sample rate, configures pyroscope."""
|
||||
mock_pyroscope = _mock_pyroscope_module()
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.get_secret_bool",
|
||||
return_value=True,
|
||||
), patch.dict(
|
||||
sys.modules,
|
||||
{"pyroscope": mock_pyroscope},
|
||||
), patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
"PYROSCOPE_APP_NAME": "myapp",
|
||||
"PYROSCOPE_SERVER_ADDRESS": "http://localhost:4040",
|
||||
"PYROSCOPE_SAMPLE_RATE": "100",
|
||||
},
|
||||
clear=False,
|
||||
):
|
||||
ProxyStartupEvent._init_pyroscope()
|
||||
mock_pyroscope.configure.assert_called_once()
|
||||
call_kw = mock_pyroscope.configure.call_args[1]
|
||||
assert call_kw["app_name"] == "myapp"
|
||||
assert call_kw["server_address"] == "http://localhost:4040"
|
||||
assert call_kw["sample_rate"] == 100
|
||||
|
||||
|
||||
def test_init_pyroscope_accepts_float_sample_rate_parsed_as_int():
|
||||
"""PYROSCOPE_SAMPLE_RATE can be a float string; it is parsed as integer."""
|
||||
mock_pyroscope = _mock_pyroscope_module()
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.get_secret_bool",
|
||||
return_value=True,
|
||||
), patch.dict(
|
||||
sys.modules,
|
||||
{"pyroscope": mock_pyroscope},
|
||||
), patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
"PYROSCOPE_APP_NAME": "myapp",
|
||||
"PYROSCOPE_SERVER_ADDRESS": "http://localhost:4040",
|
||||
"PYROSCOPE_SAMPLE_RATE": "100.7",
|
||||
},
|
||||
clear=False,
|
||||
):
|
||||
ProxyStartupEvent._init_pyroscope()
|
||||
call_kw = mock_pyroscope.configure.call_args[1]
|
||||
assert call_kw["sample_rate"] == 100
|
||||
0
tests/test_litellm/types/__init__.py
Normal file
0
tests/test_litellm/types/__init__.py
Normal file
0
tests/test_litellm/types/proxy/__init__.py
Normal file
0
tests/test_litellm/types/proxy/__init__.py
Normal file
0
tests/test_litellm/types/proxy/policy_engine/__init__.py
Normal file
0
tests/test_litellm/types/proxy/policy_engine/__init__.py
Normal file
|
|
@ -0,0 +1,152 @@
|
|||
"""
|
||||
Tests for pipeline type definitions.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.types.proxy.policy_engine.pipeline_types import (
|
||||
GuardrailPipeline,
|
||||
PipelineExecutionResult,
|
||||
PipelineStep,
|
||||
PipelineStepResult,
|
||||
)
|
||||
from litellm.types.proxy.policy_engine.policy_types import (
|
||||
Policy,
|
||||
PolicyGuardrails,
|
||||
)
|
||||
|
||||
|
||||
def test_pipeline_step_defaults():
|
||||
step = PipelineStep(guardrail="my-guard")
|
||||
assert step.on_fail == "block"
|
||||
assert step.on_pass == "allow"
|
||||
assert step.pass_data is False
|
||||
assert step.modify_response_message is None
|
||||
|
||||
|
||||
def test_pipeline_step_valid_actions():
|
||||
step = PipelineStep(guardrail="my-guard", on_fail="next", on_pass="next")
|
||||
assert step.on_fail == "next"
|
||||
assert step.on_pass == "next"
|
||||
|
||||
|
||||
def test_pipeline_step_all_action_types():
|
||||
for action in ("allow", "block", "next", "modify_response"):
|
||||
step = PipelineStep(guardrail="g", on_fail=action, on_pass=action)
|
||||
assert step.on_fail == action
|
||||
assert step.on_pass == action
|
||||
|
||||
|
||||
def test_pipeline_step_invalid_action_rejected():
|
||||
with pytest.raises(ValidationError):
|
||||
PipelineStep(guardrail="my-guard", on_fail="invalid_action")
|
||||
|
||||
|
||||
def test_pipeline_step_invalid_on_pass_rejected():
|
||||
with pytest.raises(ValidationError):
|
||||
PipelineStep(guardrail="my-guard", on_pass="skip")
|
||||
|
||||
|
||||
def test_pipeline_requires_at_least_one_step():
|
||||
with pytest.raises(ValidationError):
|
||||
GuardrailPipeline(mode="pre_call", steps=[])
|
||||
|
||||
|
||||
def test_pipeline_invalid_mode_rejected():
|
||||
with pytest.raises(ValidationError):
|
||||
GuardrailPipeline(
|
||||
mode="during_call",
|
||||
steps=[PipelineStep(guardrail="g")],
|
||||
)
|
||||
|
||||
|
||||
def test_pipeline_valid_modes():
|
||||
for mode in ("pre_call", "post_call"):
|
||||
pipeline = GuardrailPipeline(
|
||||
mode=mode,
|
||||
steps=[PipelineStep(guardrail="g")],
|
||||
)
|
||||
assert pipeline.mode == mode
|
||||
|
||||
|
||||
def test_pipeline_with_multiple_steps():
|
||||
pipeline = GuardrailPipeline(
|
||||
mode="pre_call",
|
||||
steps=[
|
||||
PipelineStep(guardrail="g1", on_fail="next", on_pass="allow"),
|
||||
PipelineStep(guardrail="g2", on_fail="block", on_pass="allow"),
|
||||
],
|
||||
)
|
||||
assert len(pipeline.steps) == 2
|
||||
assert pipeline.steps[0].guardrail == "g1"
|
||||
assert pipeline.steps[1].guardrail == "g2"
|
||||
|
||||
|
||||
def test_policy_with_pipeline_parses():
|
||||
policy = Policy(
|
||||
guardrails=PolicyGuardrails(add=["g1", "g2"]),
|
||||
pipeline=GuardrailPipeline(
|
||||
mode="pre_call",
|
||||
steps=[
|
||||
PipelineStep(guardrail="g1", on_fail="next"),
|
||||
PipelineStep(guardrail="g2"),
|
||||
],
|
||||
),
|
||||
)
|
||||
assert policy.pipeline is not None
|
||||
assert len(policy.pipeline.steps) == 2
|
||||
|
||||
|
||||
def test_policy_without_pipeline():
|
||||
policy = Policy(
|
||||
guardrails=PolicyGuardrails(add=["g1"]),
|
||||
)
|
||||
assert policy.pipeline is None
|
||||
|
||||
|
||||
def test_pipeline_step_result():
|
||||
result = PipelineStepResult(
|
||||
guardrail_name="g1",
|
||||
outcome="fail",
|
||||
action_taken="next",
|
||||
error_detail="Content policy violation",
|
||||
duration_seconds=0.05,
|
||||
)
|
||||
assert result.outcome == "fail"
|
||||
assert result.action_taken == "next"
|
||||
|
||||
|
||||
def test_pipeline_execution_result():
|
||||
result = PipelineExecutionResult(
|
||||
terminal_action="block",
|
||||
step_results=[
|
||||
PipelineStepResult(
|
||||
guardrail_name="g1",
|
||||
outcome="fail",
|
||||
action_taken="next",
|
||||
),
|
||||
PipelineStepResult(
|
||||
guardrail_name="g2",
|
||||
outcome="fail",
|
||||
action_taken="block",
|
||||
),
|
||||
],
|
||||
error_message="Content blocked",
|
||||
)
|
||||
assert result.terminal_action == "block"
|
||||
assert len(result.step_results) == 2
|
||||
|
||||
|
||||
def test_pipeline_step_extra_fields_rejected():
|
||||
with pytest.raises(ValidationError):
|
||||
PipelineStep(guardrail="g", unknown_field="value")
|
||||
|
||||
|
||||
def test_pipeline_extra_fields_rejected():
|
||||
with pytest.raises(ValidationError):
|
||||
GuardrailPipeline(
|
||||
mode="pre_call",
|
||||
steps=[PipelineStep(guardrail="g")],
|
||||
unknown="value",
|
||||
)
|
||||
|
|
@ -0,0 +1,102 @@
|
|||
"""
|
||||
Tests for pipeline field on policy CRUD types (resolver_types.py).
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.types.proxy.policy_engine.resolver_types import (
|
||||
PolicyCreateRequest,
|
||||
PolicyDBResponse,
|
||||
PolicyUpdateRequest,
|
||||
)
|
||||
|
||||
|
||||
def test_policy_create_request_with_pipeline():
|
||||
pipeline_data = {
|
||||
"mode": "pre_call",
|
||||
"steps": [
|
||||
{"guardrail": "g1", "on_fail": "next", "on_pass": "allow"},
|
||||
{"guardrail": "g2", "on_fail": "block", "on_pass": "allow"},
|
||||
],
|
||||
}
|
||||
req = PolicyCreateRequest(
|
||||
policy_name="test-policy",
|
||||
guardrails_add=["g1", "g2"],
|
||||
pipeline=pipeline_data,
|
||||
)
|
||||
assert req.pipeline is not None
|
||||
assert req.pipeline["mode"] == "pre_call"
|
||||
assert len(req.pipeline["steps"]) == 2
|
||||
|
||||
|
||||
def test_policy_create_request_without_pipeline():
|
||||
req = PolicyCreateRequest(
|
||||
policy_name="test-policy",
|
||||
guardrails_add=["g1"],
|
||||
)
|
||||
assert req.pipeline is None
|
||||
|
||||
|
||||
def test_policy_update_request_with_pipeline():
|
||||
pipeline_data = {
|
||||
"mode": "pre_call",
|
||||
"steps": [
|
||||
{"guardrail": "g1", "on_fail": "block", "on_pass": "allow"},
|
||||
],
|
||||
}
|
||||
req = PolicyUpdateRequest(pipeline=pipeline_data)
|
||||
assert req.pipeline is not None
|
||||
assert req.pipeline["steps"][0]["guardrail"] == "g1"
|
||||
|
||||
|
||||
def test_policy_db_response_with_pipeline():
|
||||
pipeline_data = {
|
||||
"mode": "pre_call",
|
||||
"steps": [
|
||||
{"guardrail": "g1", "on_fail": "next", "on_pass": "allow"},
|
||||
{"guardrail": "g2", "on_fail": "block", "on_pass": "allow"},
|
||||
],
|
||||
}
|
||||
resp = PolicyDBResponse(
|
||||
policy_id="test-id",
|
||||
policy_name="test-policy",
|
||||
guardrails_add=["g1", "g2"],
|
||||
pipeline=pipeline_data,
|
||||
)
|
||||
assert resp.pipeline is not None
|
||||
assert resp.pipeline["mode"] == "pre_call"
|
||||
dumped = resp.model_dump()
|
||||
assert dumped["pipeline"]["steps"][0]["guardrail"] == "g1"
|
||||
|
||||
|
||||
def test_policy_db_response_without_pipeline():
|
||||
resp = PolicyDBResponse(
|
||||
policy_id="test-id",
|
||||
policy_name="test-policy",
|
||||
)
|
||||
assert resp.pipeline is None
|
||||
dumped = resp.model_dump()
|
||||
assert dumped["pipeline"] is None
|
||||
|
||||
|
||||
def test_policy_create_request_roundtrip():
|
||||
pipeline_data = {
|
||||
"mode": "post_call",
|
||||
"steps": [
|
||||
{
|
||||
"guardrail": "g1",
|
||||
"on_fail": "modify_response",
|
||||
"on_pass": "next",
|
||||
"pass_data": True,
|
||||
"modify_response_message": "custom msg",
|
||||
},
|
||||
],
|
||||
}
|
||||
req = PolicyCreateRequest(
|
||||
policy_name="roundtrip-test",
|
||||
guardrails_add=["g1"],
|
||||
pipeline=pipeline_data,
|
||||
)
|
||||
dumped = req.model_dump()
|
||||
restored = PolicyCreateRequest(**dumped)
|
||||
assert restored.pipeline == pipeline_data
|
||||
|
|
@ -8,7 +8,7 @@ async function globalSetup() {
|
|||
await page.goto("http://localhost:4000/ui/login");
|
||||
await page.getByPlaceholder("Enter your username").fill(users[Role.ProxyAdmin].email);
|
||||
await page.getByPlaceholder("Enter your password").fill(users[Role.ProxyAdmin].password);
|
||||
const loginButton = page.getByRole("button", { name: "Login" });
|
||||
const loginButton = page.getByRole("button", { name: "Login", exact: true });
|
||||
await loginButton.click();
|
||||
await page.waitForSelector("text=AI Gateway");
|
||||
await page.context().storageState({ path: "admin.storageState.json" });
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ test("user can log in", async ({ page }) => {
|
|||
await page.goto("http://localhost:4000/ui/login");
|
||||
await page.getByPlaceholder("Enter your username").fill(users[Role.ProxyAdmin].email);
|
||||
await page.getByPlaceholder("Enter your password").fill(users[Role.ProxyAdmin].password);
|
||||
const loginButton = page.getByRole("button", { name: "Login" });
|
||||
const loginButton = page.getByRole("button", { name: "Login", exact: true });
|
||||
await expect(loginButton).toBeEnabled();
|
||||
await loginButton.click();
|
||||
await expect(page.getByText("AI Gateway")).toBeVisible();
|
||||
|
|
|
|||
|
|
@ -0,0 +1,63 @@
|
|||
import { useQuery, useQueryClient } from "@tanstack/react-query";
|
||||
import {
|
||||
getProxyBaseUrl,
|
||||
getGlobalLitellmHeaderName,
|
||||
deriveErrorMessage,
|
||||
handleError,
|
||||
} from "@/components/networking";
|
||||
import { all_admin_roles } from "@/utils/roles";
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
import { AccessGroupResponse, accessGroupKeys } from "./useAccessGroups";
|
||||
|
||||
// ── Fetch function ───────────────────────────────────────────────────────────
|
||||
|
||||
const fetchAccessGroupDetails = async (
|
||||
accessToken: string,
|
||||
accessGroupId: string,
|
||||
): Promise<AccessGroupResponse> => {
|
||||
const baseUrl = getProxyBaseUrl();
|
||||
const url = `${baseUrl}/v1/access_group/${encodeURIComponent(accessGroupId)}`;
|
||||
|
||||
const response = await fetch(url, {
|
||||
method: "GET",
|
||||
headers: {
|
||||
[getGlobalLitellmHeaderName()]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorData = await response.json();
|
||||
const errorMessage = deriveErrorMessage(errorData);
|
||||
handleError(errorMessage);
|
||||
throw new Error(errorMessage);
|
||||
}
|
||||
|
||||
return response.json();
|
||||
};
|
||||
|
||||
// ── Hook ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
export const useAccessGroupDetails = (accessGroupId?: string) => {
|
||||
const { accessToken, userRole } = useAuthorized();
|
||||
const queryClient = useQueryClient();
|
||||
|
||||
return useQuery<AccessGroupResponse>({
|
||||
queryKey: accessGroupKeys.detail(accessGroupId!),
|
||||
queryFn: async () => fetchAccessGroupDetails(accessToken!, accessGroupId!),
|
||||
enabled:
|
||||
Boolean(accessToken && accessGroupId) &&
|
||||
all_admin_roles.includes(userRole || ""),
|
||||
|
||||
// Seed from the list cache when available
|
||||
initialData: () => {
|
||||
if (!accessGroupId) return undefined;
|
||||
|
||||
const groups = queryClient.getQueryData<AccessGroupResponse[]>(
|
||||
accessGroupKeys.list({}),
|
||||
);
|
||||
|
||||
return groups?.find((g) => g.access_group_id === accessGroupId);
|
||||
},
|
||||
});
|
||||
};
|
||||
|
|
@ -0,0 +1,242 @@
|
|||
/* @vitest-environment jsdom */
|
||||
import React from "react";
|
||||
import { renderHook, waitFor } from "@testing-library/react";
|
||||
import { describe, it, expect, vi, beforeEach } from "vitest";
|
||||
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
|
||||
import { useAccessGroups, AccessGroupResponse } from "./useAccessGroups";
|
||||
import * as networking from "@/components/networking";
|
||||
|
||||
vi.mock("@/components/networking", () => ({
|
||||
getProxyBaseUrl: vi.fn(() => "http://proxy.example"),
|
||||
getGlobalLitellmHeaderName: vi.fn(() => "Authorization"),
|
||||
deriveErrorMessage: vi.fn((data: unknown) => (data as { detail?: string })?.detail ?? "Unknown error"),
|
||||
handleError: vi.fn(),
|
||||
}));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
|
||||
default: vi.fn(() => ({
|
||||
accessToken: "test-token-123",
|
||||
userRole: "Admin",
|
||||
})),
|
||||
}));
|
||||
|
||||
const createQueryClient = () =>
|
||||
new QueryClient({
|
||||
defaultOptions: {
|
||||
queries: {
|
||||
retry: false,
|
||||
gcTime: 0,
|
||||
},
|
||||
},
|
||||
});
|
||||
|
||||
const wrapper = ({ children }: { children: React.ReactNode }) => {
|
||||
const queryClient = createQueryClient();
|
||||
return React.createElement(QueryClientProvider, { client: queryClient }, children);
|
||||
};
|
||||
|
||||
const mockAccessToken = "test-token-123";
|
||||
const mockAccessGroups: AccessGroupResponse[] = [
|
||||
{
|
||||
access_group_id: "ag-1",
|
||||
access_group_name: "Group One",
|
||||
description: "First group",
|
||||
access_model_ids: [],
|
||||
access_mcp_server_ids: [],
|
||||
access_agent_ids: [],
|
||||
assigned_team_ids: [],
|
||||
assigned_key_ids: [],
|
||||
created_at: "2025-01-01T00:00:00Z",
|
||||
created_by: "user-1",
|
||||
updated_at: "2025-01-01T00:00:00Z",
|
||||
updated_by: "user-1",
|
||||
},
|
||||
];
|
||||
|
||||
const fetchMock = vi.fn();
|
||||
|
||||
describe("useAccessGroups", () => {
|
||||
beforeEach(async () => {
|
||||
vi.clearAllMocks();
|
||||
vi.mocked(networking.getProxyBaseUrl).mockReturnValue("http://proxy.example");
|
||||
vi.mocked(networking.getGlobalLitellmHeaderName).mockReturnValue("Authorization");
|
||||
|
||||
const useAuthorizedModule = await import("@/app/(dashboard)/hooks/useAuthorized");
|
||||
vi.mocked(useAuthorizedModule.default).mockReturnValue({
|
||||
accessToken: mockAccessToken,
|
||||
userRole: "Admin",
|
||||
} as any);
|
||||
|
||||
global.fetch = fetchMock;
|
||||
});
|
||||
|
||||
it("should return hook result without errors", () => {
|
||||
fetchMock.mockResolvedValue({
|
||||
ok: true,
|
||||
json: () => Promise.resolve([]),
|
||||
} as Response);
|
||||
|
||||
const { result } = renderHook(() => useAccessGroups(), { wrapper });
|
||||
|
||||
expect(result.current).toBeDefined();
|
||||
expect(result.current).toHaveProperty("data");
|
||||
expect(result.current).toHaveProperty("isSuccess");
|
||||
expect(result.current).toHaveProperty("isError");
|
||||
expect(result.current).toHaveProperty("status");
|
||||
});
|
||||
|
||||
it("should return access groups when access token and admin role are present", async () => {
|
||||
fetchMock.mockResolvedValue({
|
||||
ok: true,
|
||||
json: () => Promise.resolve(mockAccessGroups),
|
||||
} as Response);
|
||||
|
||||
const { result } = renderHook(() => useAccessGroups(), { wrapper });
|
||||
|
||||
await waitFor(() => {
|
||||
expect(result.current.isSuccess).toBe(true);
|
||||
});
|
||||
|
||||
expect(fetchMock).toHaveBeenCalledWith(
|
||||
"http://proxy.example/v1/access_group",
|
||||
expect.objectContaining({
|
||||
method: "GET",
|
||||
headers: expect.objectContaining({
|
||||
Authorization: `Bearer ${mockAccessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
}),
|
||||
}),
|
||||
);
|
||||
expect(result.current.data).toEqual(mockAccessGroups);
|
||||
});
|
||||
|
||||
it("should not fetch when access token is null", async () => {
|
||||
const useAuthorizedModule = await import("@/app/(dashboard)/hooks/useAuthorized");
|
||||
vi.mocked(useAuthorizedModule.default).mockReturnValue({
|
||||
accessToken: null,
|
||||
userRole: "Admin",
|
||||
} as any);
|
||||
|
||||
const { result } = renderHook(() => useAccessGroups(), { wrapper });
|
||||
|
||||
expect(result.current.isFetching).toBe(false);
|
||||
expect(result.current.isLoading).toBe(false);
|
||||
expect(result.current.data).toBeUndefined();
|
||||
expect(fetchMock).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should not fetch when access token is empty string", async () => {
|
||||
const useAuthorizedModule = await import("@/app/(dashboard)/hooks/useAuthorized");
|
||||
vi.mocked(useAuthorizedModule.default).mockReturnValue({
|
||||
accessToken: "",
|
||||
userRole: "Admin",
|
||||
} as any);
|
||||
|
||||
const { result } = renderHook(() => useAccessGroups(), { wrapper });
|
||||
|
||||
expect(result.current.isFetching).toBe(false);
|
||||
expect(result.current.isLoading).toBe(false);
|
||||
expect(result.current.data).toBeUndefined();
|
||||
expect(fetchMock).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should not fetch when user role is not an admin role", async () => {
|
||||
const useAuthorizedModule = await import("@/app/(dashboard)/hooks/useAuthorized");
|
||||
vi.mocked(useAuthorizedModule.default).mockReturnValue({
|
||||
accessToken: mockAccessToken,
|
||||
userRole: "Viewer",
|
||||
} as any);
|
||||
|
||||
const { result } = renderHook(() => useAccessGroups(), { wrapper });
|
||||
|
||||
expect(result.current.isFetching).toBe(false);
|
||||
expect(result.current.isLoading).toBe(false);
|
||||
expect(result.current.data).toBeUndefined();
|
||||
expect(fetchMock).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should not fetch when user role is null", async () => {
|
||||
const useAuthorizedModule = await import("@/app/(dashboard)/hooks/useAuthorized");
|
||||
vi.mocked(useAuthorizedModule.default).mockReturnValue({
|
||||
accessToken: mockAccessToken,
|
||||
userRole: null,
|
||||
} as any);
|
||||
|
||||
const { result } = renderHook(() => useAccessGroups(), { wrapper });
|
||||
|
||||
expect(result.current.isFetching).toBe(false);
|
||||
expect(result.current.isLoading).toBe(false);
|
||||
expect(result.current.data).toBeUndefined();
|
||||
expect(fetchMock).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should fetch when user role is proxy_admin", async () => {
|
||||
const useAuthorizedModule = await import("@/app/(dashboard)/hooks/useAuthorized");
|
||||
vi.mocked(useAuthorizedModule.default).mockReturnValue({
|
||||
accessToken: mockAccessToken,
|
||||
userRole: "proxy_admin",
|
||||
} as any);
|
||||
|
||||
fetchMock.mockResolvedValue({
|
||||
ok: true,
|
||||
json: () => Promise.resolve(mockAccessGroups),
|
||||
} as Response);
|
||||
|
||||
const { result } = renderHook(() => useAccessGroups(), { wrapper });
|
||||
|
||||
await waitFor(() => {
|
||||
expect(result.current.isSuccess).toBe(true);
|
||||
});
|
||||
|
||||
expect(fetchMock).toHaveBeenCalled();
|
||||
expect(result.current.data).toEqual(mockAccessGroups);
|
||||
});
|
||||
|
||||
it("should expose error state when fetch fails", async () => {
|
||||
fetchMock.mockResolvedValue({
|
||||
ok: false,
|
||||
json: () => Promise.resolve({ detail: "Forbidden" }),
|
||||
} as Response);
|
||||
vi.mocked(networking.deriveErrorMessage).mockReturnValue("Forbidden");
|
||||
|
||||
const { result } = renderHook(() => useAccessGroups(), { wrapper });
|
||||
|
||||
await waitFor(() => {
|
||||
expect(result.current.isError).toBe(true);
|
||||
});
|
||||
|
||||
expect(result.current.error).toBeInstanceOf(Error);
|
||||
expect((result.current.error as Error).message).toBe("Forbidden");
|
||||
expect(result.current.data).toBeUndefined();
|
||||
expect(networking.handleError).toHaveBeenCalledWith("Forbidden");
|
||||
});
|
||||
|
||||
it("should return empty array when API returns empty list", async () => {
|
||||
fetchMock.mockResolvedValue({
|
||||
ok: true,
|
||||
json: () => Promise.resolve([]),
|
||||
} as Response);
|
||||
|
||||
const { result } = renderHook(() => useAccessGroups(), { wrapper });
|
||||
|
||||
await waitFor(() => {
|
||||
expect(result.current.isSuccess).toBe(true);
|
||||
});
|
||||
|
||||
expect(result.current.data).toEqual([]);
|
||||
});
|
||||
|
||||
it("should propagate network errors", async () => {
|
||||
const networkError = new Error("Network failure");
|
||||
fetchMock.mockRejectedValue(networkError);
|
||||
|
||||
const { result } = renderHook(() => useAccessGroups(), { wrapper });
|
||||
|
||||
await waitFor(() => {
|
||||
expect(result.current.isError).toBe(true);
|
||||
});
|
||||
|
||||
expect(result.current.error).toEqual(networkError);
|
||||
expect(result.current.data).toBeUndefined();
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,70 @@
|
|||
import { useQuery } from "@tanstack/react-query";
|
||||
import { createQueryKeys } from "../common/queryKeysFactory";
|
||||
import {
|
||||
getProxyBaseUrl,
|
||||
getGlobalLitellmHeaderName,
|
||||
deriveErrorMessage,
|
||||
handleError,
|
||||
} from "@/components/networking";
|
||||
import { all_admin_roles } from "@/utils/roles";
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
|
||||
// ── Types ────────────────────────────────────────────────────────────────────
|
||||
|
||||
export interface AccessGroupResponse {
|
||||
access_group_id: string;
|
||||
access_group_name: string;
|
||||
description: string | null;
|
||||
access_model_ids: string[];
|
||||
access_mcp_server_ids: string[];
|
||||
access_agent_ids: string[];
|
||||
assigned_team_ids: string[];
|
||||
assigned_key_ids: string[];
|
||||
created_at: string;
|
||||
created_by: string | null;
|
||||
updated_at: string;
|
||||
updated_by: string | null;
|
||||
}
|
||||
|
||||
// ── Query keys (shared across access-group hooks) ────────────────────────────
|
||||
|
||||
export const accessGroupKeys = createQueryKeys("accessGroups");
|
||||
|
||||
// ── Fetch function ───────────────────────────────────────────────────────────
|
||||
|
||||
const fetchAccessGroups = async (
|
||||
accessToken: string,
|
||||
): Promise<AccessGroupResponse[]> => {
|
||||
const baseUrl = getProxyBaseUrl();
|
||||
const url = `${baseUrl}/v1/access_group`;
|
||||
|
||||
const response = await fetch(url, {
|
||||
method: "GET",
|
||||
headers: {
|
||||
[getGlobalLitellmHeaderName()]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorData = await response.json();
|
||||
const errorMessage = deriveErrorMessage(errorData);
|
||||
handleError(errorMessage);
|
||||
throw new Error(errorMessage);
|
||||
}
|
||||
|
||||
return response.json();
|
||||
};
|
||||
|
||||
// ── Hook ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
export const useAccessGroups = () => {
|
||||
const { accessToken, userRole } = useAuthorized();
|
||||
|
||||
return useQuery<AccessGroupResponse[]>({
|
||||
queryKey: accessGroupKeys.list({}),
|
||||
queryFn: async () => fetchAccessGroups(accessToken!),
|
||||
enabled:
|
||||
Boolean(accessToken) && all_admin_roles.includes(userRole || ""),
|
||||
});
|
||||
};
|
||||
|
|
@ -0,0 +1,68 @@
|
|||
import { useMutation, useQueryClient } from "@tanstack/react-query";
|
||||
import {
|
||||
getProxyBaseUrl,
|
||||
getGlobalLitellmHeaderName,
|
||||
deriveErrorMessage,
|
||||
handleError,
|
||||
} from "@/components/networking";
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
import { AccessGroupResponse, accessGroupKeys } from "./useAccessGroups";
|
||||
|
||||
// ── Types ────────────────────────────────────────────────────────────────────
|
||||
|
||||
export interface AccessGroupCreateParams {
|
||||
access_group_name: string;
|
||||
description?: string | null;
|
||||
access_model_ids?: string[];
|
||||
access_mcp_server_ids?: string[];
|
||||
access_agent_ids?: string[];
|
||||
assigned_team_ids?: string[];
|
||||
assigned_key_ids?: string[];
|
||||
}
|
||||
|
||||
// ── Fetch function ───────────────────────────────────────────────────────────
|
||||
|
||||
const createAccessGroup = async (
|
||||
accessToken: string,
|
||||
params: AccessGroupCreateParams,
|
||||
): Promise<AccessGroupResponse> => {
|
||||
const baseUrl = getProxyBaseUrl();
|
||||
const url = `${baseUrl}/v1/access_group`;
|
||||
|
||||
const response = await fetch(url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
[getGlobalLitellmHeaderName()]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
body: JSON.stringify(params),
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorData = await response.json();
|
||||
const errorMessage = deriveErrorMessage(errorData);
|
||||
handleError(errorMessage);
|
||||
throw new Error(errorMessage);
|
||||
}
|
||||
|
||||
return response.json();
|
||||
};
|
||||
|
||||
// ── Hook ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
export const useCreateAccessGroup = () => {
|
||||
const { accessToken } = useAuthorized();
|
||||
const queryClient = useQueryClient();
|
||||
|
||||
return useMutation<AccessGroupResponse, Error, AccessGroupCreateParams>({
|
||||
mutationFn: async (params) => {
|
||||
if (!accessToken) {
|
||||
throw new Error("Access token is required");
|
||||
}
|
||||
return createAccessGroup(accessToken, params);
|
||||
},
|
||||
onSuccess: () => {
|
||||
queryClient.invalidateQueries({ queryKey: accessGroupKeys.all });
|
||||
},
|
||||
});
|
||||
};
|
||||
|
|
@ -0,0 +1,55 @@
|
|||
import { useMutation, useQueryClient } from "@tanstack/react-query";
|
||||
import {
|
||||
getProxyBaseUrl,
|
||||
getGlobalLitellmHeaderName,
|
||||
deriveErrorMessage,
|
||||
handleError,
|
||||
} from "@/components/networking";
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
import { accessGroupKeys } from "./useAccessGroups";
|
||||
|
||||
// ── Fetch function ───────────────────────────────────────────────────────────
|
||||
|
||||
const deleteAccessGroup = async (
|
||||
accessToken: string,
|
||||
accessGroupId: string,
|
||||
): Promise<void> => {
|
||||
const baseUrl = getProxyBaseUrl();
|
||||
const url = `${baseUrl}/v1/access_group/${encodeURIComponent(accessGroupId)}`;
|
||||
|
||||
const response = await fetch(url, {
|
||||
method: "DELETE",
|
||||
headers: {
|
||||
[getGlobalLitellmHeaderName()]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorData = await response.json();
|
||||
const errorMessage = deriveErrorMessage(errorData);
|
||||
handleError(errorMessage);
|
||||
throw new Error(errorMessage);
|
||||
}
|
||||
|
||||
// 204 No Content — nothing to parse
|
||||
};
|
||||
|
||||
// ── Hook ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
export const useDeleteAccessGroup = () => {
|
||||
const { accessToken } = useAuthorized();
|
||||
const queryClient = useQueryClient();
|
||||
|
||||
return useMutation<void, Error, string>({
|
||||
mutationFn: async (accessGroupId) => {
|
||||
if (!accessToken) {
|
||||
throw new Error("Access token is required");
|
||||
}
|
||||
return deleteAccessGroup(accessToken, accessGroupId);
|
||||
},
|
||||
onSuccess: () => {
|
||||
queryClient.invalidateQueries({ queryKey: accessGroupKeys.all });
|
||||
},
|
||||
});
|
||||
};
|
||||
|
|
@ -0,0 +1,77 @@
|
|||
import { useMutation, useQueryClient } from "@tanstack/react-query";
|
||||
import {
|
||||
getProxyBaseUrl,
|
||||
getGlobalLitellmHeaderName,
|
||||
deriveErrorMessage,
|
||||
handleError,
|
||||
} from "@/components/networking";
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
import { AccessGroupResponse, accessGroupKeys } from "./useAccessGroups";
|
||||
|
||||
// ── Types ────────────────────────────────────────────────────────────────────
|
||||
|
||||
export interface AccessGroupUpdateParams {
|
||||
access_group_name?: string;
|
||||
description?: string | null;
|
||||
access_model_ids?: string[];
|
||||
access_mcp_server_ids?: string[];
|
||||
access_agent_ids?: string[];
|
||||
assigned_team_ids?: string[];
|
||||
assigned_key_ids?: string[];
|
||||
}
|
||||
|
||||
export interface EditAccessGroupVariables {
|
||||
accessGroupId: string;
|
||||
params: AccessGroupUpdateParams;
|
||||
}
|
||||
|
||||
// ── Fetch function ───────────────────────────────────────────────────────────
|
||||
|
||||
const updateAccessGroup = async (
|
||||
accessToken: string,
|
||||
accessGroupId: string,
|
||||
params: AccessGroupUpdateParams,
|
||||
): Promise<AccessGroupResponse> => {
|
||||
const baseUrl = getProxyBaseUrl();
|
||||
const url = `${baseUrl}/v1/access_group/${encodeURIComponent(accessGroupId)}`;
|
||||
|
||||
const response = await fetch(url, {
|
||||
method: "PUT",
|
||||
headers: {
|
||||
[getGlobalLitellmHeaderName()]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
body: JSON.stringify(params),
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorData = await response.json();
|
||||
const errorMessage = deriveErrorMessage(errorData);
|
||||
handleError(errorMessage);
|
||||
throw new Error(errorMessage);
|
||||
}
|
||||
|
||||
return response.json();
|
||||
};
|
||||
|
||||
// ── Hook ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
export const useEditAccessGroup = () => {
|
||||
const { accessToken } = useAuthorized();
|
||||
const queryClient = useQueryClient();
|
||||
|
||||
return useMutation<AccessGroupResponse, Error, EditAccessGroupVariables>({
|
||||
mutationFn: async ({ accessGroupId, params }) => {
|
||||
if (!accessToken) {
|
||||
throw new Error("Access token is required");
|
||||
}
|
||||
return updateAccessGroup(accessToken, accessGroupId, params);
|
||||
},
|
||||
onSuccess: (_data, { accessGroupId }) => {
|
||||
queryClient.invalidateQueries({ queryKey: accessGroupKeys.all });
|
||||
queryClient.invalidateQueries({
|
||||
queryKey: accessGroupKeys.detail(accessGroupId),
|
||||
});
|
||||
},
|
||||
});
|
||||
};
|
||||
|
|
@ -35,6 +35,7 @@ import TransformRequestPanel from "@/components/transform_request";
|
|||
import UIThemeSettings from "@/components/ui_theme_settings";
|
||||
import Usage from "@/components/usage";
|
||||
import UserDashboard from "@/components/user_dashboard";
|
||||
import { AccessGroupsPage } from "@/components/AccessGroups/AccessGroupsPage";
|
||||
import VectorStoreManagement from "@/components/vector_store_management";
|
||||
import SpendLogsTable from "@/components/view_logs";
|
||||
import ViewUserDashboard from "@/components/view_users";
|
||||
|
|
@ -542,6 +543,8 @@ function CreateKeyPageContent() {
|
|||
<TagManagement accessToken={accessToken} userRole={userRole} userID={userID} />
|
||||
) : page == "claude-code-plugins" ? (
|
||||
<ClaudeCodePluginsPanel accessToken={accessToken} userRole={userRole} />
|
||||
) : page == "access-groups" ? (
|
||||
<AccessGroupsPage />
|
||||
) : page == "vector-stores" ? (
|
||||
<VectorStoreManagement accessToken={accessToken} userRole={userRole} userID={userID} />
|
||||
) : page == "new_usage" ? (
|
||||
|
|
|
|||
|
|
@ -0,0 +1,384 @@
|
|||
import { useAccessGroupDetails } from "@/app/(dashboard)/hooks/accessGroups/useAccessGroupDetails";
|
||||
import { AccessGroupResponse } from "@/app/(dashboard)/hooks/accessGroups/useAccessGroups";
|
||||
import { screen } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import { renderWithProviders } from "../../../tests/test-utils";
|
||||
import { AccessGroupDetail } from "./AccessGroupsDetailsPage";
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/accessGroups/useAccessGroupDetails");
|
||||
vi.mock("./AccessGroupsModal/AccessGroupEditModal", () => ({
|
||||
AccessGroupEditModal: ({
|
||||
visible,
|
||||
onCancel,
|
||||
}: {
|
||||
visible: boolean;
|
||||
onCancel: () => void;
|
||||
}) =>
|
||||
visible ? (
|
||||
<div role="dialog" aria-label="Edit Access Group">
|
||||
<button onClick={onCancel}>Close Modal</button>
|
||||
</div>
|
||||
) : null,
|
||||
}));
|
||||
|
||||
const mockUseAccessGroupDetails = vi.mocked(useAccessGroupDetails);
|
||||
|
||||
const baseMockReturnValue = {
|
||||
data: undefined,
|
||||
isLoading: false,
|
||||
isError: false,
|
||||
error: null,
|
||||
isFetching: false,
|
||||
isPending: false,
|
||||
isSuccess: true,
|
||||
status: "success" as const,
|
||||
dataUpdatedAt: 0,
|
||||
errorUpdatedAt: 0,
|
||||
failureCount: 0,
|
||||
failureReason: null,
|
||||
errorUpdateCount: 0,
|
||||
isFetched: true,
|
||||
isFetchedAfterMount: true,
|
||||
isRefetching: false,
|
||||
isLoadingError: false,
|
||||
isPaused: false,
|
||||
isPlaceholderData: false,
|
||||
isRefetchError: false,
|
||||
isStale: false,
|
||||
fetchStatus: "idle" as const,
|
||||
refetch: vi.fn(),
|
||||
} as unknown as ReturnType<typeof useAccessGroupDetails>;
|
||||
|
||||
const createMockAccessGroup = (
|
||||
overrides: Partial<AccessGroupResponse> = {}
|
||||
): AccessGroupResponse => ({
|
||||
access_group_id: "ag-1",
|
||||
access_group_name: "Test Group",
|
||||
description: "A test access group",
|
||||
access_model_ids: ["model-1", "model-2"],
|
||||
access_mcp_server_ids: ["mcp-1"],
|
||||
access_agent_ids: ["agent-1"],
|
||||
assigned_team_ids: ["team-1"],
|
||||
assigned_key_ids: ["key-1", "key-2"],
|
||||
created_at: "2025-01-01T00:00:00Z",
|
||||
created_by: null,
|
||||
updated_at: "2025-01-02T00:00:00Z",
|
||||
updated_by: null,
|
||||
...overrides,
|
||||
});
|
||||
|
||||
describe("AccessGroupDetail", () => {
|
||||
const mockOnBack = vi.fn();
|
||||
const accessGroupId = "ag-1";
|
||||
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
mockUseAccessGroupDetails.mockReturnValue({
|
||||
...baseMockReturnValue,
|
||||
data: createMockAccessGroup(),
|
||||
} as ReturnType<typeof useAccessGroupDetails>);
|
||||
});
|
||||
|
||||
it("should render the component", () => {
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
expect(screen.getByRole("heading", { name: "Test Group" })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should not show access group content when loading", () => {
|
||||
mockUseAccessGroupDetails.mockReturnValue({
|
||||
...baseMockReturnValue,
|
||||
data: undefined,
|
||||
isLoading: true,
|
||||
} as ReturnType<typeof useAccessGroupDetails>);
|
||||
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
expect(screen.queryByRole("heading", { name: "Test Group" })).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show empty state when access group is not found", () => {
|
||||
mockUseAccessGroupDetails.mockReturnValue({
|
||||
...baseMockReturnValue,
|
||||
data: undefined,
|
||||
isLoading: false,
|
||||
} as ReturnType<typeof useAccessGroupDetails>);
|
||||
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
expect(screen.getByText("Access group not found")).toBeInTheDocument();
|
||||
expect(screen.getByRole("button")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should call onBack when back button is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
const buttons = screen.getAllByRole("button");
|
||||
const backButton = buttons.find((btn) => !btn.textContent?.includes("Edit"));
|
||||
await user.click(backButton!);
|
||||
|
||||
expect(mockOnBack).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
it("should display access group name and ID", () => {
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
expect(screen.getByRole("heading", { name: "Test Group" })).toBeInTheDocument();
|
||||
expect(screen.getByText(/ID:/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display description in Group Details", () => {
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
expect(screen.getByText("Group Details")).toBeInTheDocument();
|
||||
expect(screen.getByText("A test access group")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display em dash when description is empty", () => {
|
||||
mockUseAccessGroupDetails.mockReturnValue({
|
||||
...baseMockReturnValue,
|
||||
data: createMockAccessGroup({ description: null }),
|
||||
} as ReturnType<typeof useAccessGroupDetails>);
|
||||
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
expect(screen.getByText("—")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should open edit modal when Edit Access Group button is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
expect(screen.queryByRole("dialog", { name: "Edit Access Group" })).not.toBeInTheDocument();
|
||||
|
||||
const editButton = screen.getByRole("button", { name: /Edit Access Group/i });
|
||||
await user.click(editButton);
|
||||
|
||||
expect(screen.getByRole("dialog", { name: "Edit Access Group" })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should close edit modal when Close Modal is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /Edit Access Group/i }));
|
||||
expect(screen.getByRole("dialog", { name: "Edit Access Group" })).toBeInTheDocument();
|
||||
|
||||
await user.click(screen.getByRole("button", { name: "Close Modal" }));
|
||||
expect(screen.queryByRole("dialog", { name: "Edit Access Group" })).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display attached keys", () => {
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
expect(screen.getByText("Attached Keys")).toBeInTheDocument();
|
||||
expect(screen.getByText("key-1")).toBeInTheDocument();
|
||||
expect(screen.getByText("key-2")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display attached teams", () => {
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
expect(screen.getByText("Attached Teams")).toBeInTheDocument();
|
||||
expect(screen.getByText("team-1")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show View All button for keys when more than 5", () => {
|
||||
mockUseAccessGroupDetails.mockReturnValue({
|
||||
...baseMockReturnValue,
|
||||
data: createMockAccessGroup({
|
||||
assigned_key_ids: ["k1", "k2", "k3", "k4", "k5", "k6"],
|
||||
}),
|
||||
} as ReturnType<typeof useAccessGroupDetails>);
|
||||
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
expect(screen.getByRole("button", { name: "View All (6)" })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should toggle between View All and Show Less for keys", async () => {
|
||||
const user = userEvent.setup();
|
||||
mockUseAccessGroupDetails.mockReturnValue({
|
||||
...baseMockReturnValue,
|
||||
data: createMockAccessGroup({
|
||||
assigned_key_ids: ["k1", "k2", "k3", "k4", "k5", "k6"],
|
||||
}),
|
||||
} as ReturnType<typeof useAccessGroupDetails>);
|
||||
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: "View All (6)" }));
|
||||
expect(screen.getByRole("button", { name: "Show Less" })).toBeInTheDocument();
|
||||
|
||||
await user.click(screen.getByRole("button", { name: "Show Less" }));
|
||||
expect(screen.getByRole("button", { name: "View All (6)" })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show View All button for teams when more than 5", () => {
|
||||
mockUseAccessGroupDetails.mockReturnValue({
|
||||
...baseMockReturnValue,
|
||||
data: createMockAccessGroup({
|
||||
assigned_team_ids: ["t1", "t2", "t3", "t4", "t5", "t6"],
|
||||
}),
|
||||
} as ReturnType<typeof useAccessGroupDetails>);
|
||||
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
expect(screen.getByRole("button", { name: "View All (6)" })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show empty state when no keys attached", () => {
|
||||
mockUseAccessGroupDetails.mockReturnValue({
|
||||
...baseMockReturnValue,
|
||||
data: createMockAccessGroup({ assigned_key_ids: [] }),
|
||||
} as ReturnType<typeof useAccessGroupDetails>);
|
||||
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
expect(screen.getByText("No keys attached")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show empty state when no teams attached", () => {
|
||||
mockUseAccessGroupDetails.mockReturnValue({
|
||||
...baseMockReturnValue,
|
||||
data: createMockAccessGroup({ assigned_team_ids: [] }),
|
||||
} as ReturnType<typeof useAccessGroupDetails>);
|
||||
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
expect(screen.getByText("No teams attached")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display Models tab with model IDs", () => {
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
expect(screen.getByRole("tab", { name: /Models/i })).toBeInTheDocument();
|
||||
expect(screen.getByText("model-1")).toBeInTheDocument();
|
||||
expect(screen.getByText("model-2")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display MCP Servers tab with server IDs", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
const mcpTab = screen.getByRole("tab", { name: /MCP Servers/i });
|
||||
expect(mcpTab).toBeInTheDocument();
|
||||
await user.click(mcpTab);
|
||||
expect(screen.getByText("mcp-1")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display Agents tab with agent IDs", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
const agentsTab = screen.getByRole("tab", { name: /Agents/i });
|
||||
expect(agentsTab).toBeInTheDocument();
|
||||
await user.click(agentsTab);
|
||||
expect(screen.getByText("agent-1")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show empty state in Models tab when no models assigned", () => {
|
||||
mockUseAccessGroupDetails.mockReturnValue({
|
||||
...baseMockReturnValue,
|
||||
data: createMockAccessGroup({ access_model_ids: [] }),
|
||||
} as ReturnType<typeof useAccessGroupDetails>);
|
||||
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
expect(screen.getByText("No models assigned to this group")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show empty state in MCP Servers tab when none assigned", async () => {
|
||||
const user = userEvent.setup();
|
||||
mockUseAccessGroupDetails.mockReturnValue({
|
||||
...baseMockReturnValue,
|
||||
data: createMockAccessGroup({ access_mcp_server_ids: [] }),
|
||||
} as ReturnType<typeof useAccessGroupDetails>);
|
||||
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
await user.click(screen.getByRole("tab", { name: /MCP Servers/i }));
|
||||
expect(screen.getByText("No MCP servers assigned to this group")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show empty state in Agents tab when none assigned", async () => {
|
||||
const user = userEvent.setup();
|
||||
mockUseAccessGroupDetails.mockReturnValue({
|
||||
...baseMockReturnValue,
|
||||
data: createMockAccessGroup({ access_agent_ids: [] }),
|
||||
} as ReturnType<typeof useAccessGroupDetails>);
|
||||
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
await user.click(screen.getByRole("tab", { name: /Agents/i }));
|
||||
expect(screen.getByText("No agents assigned to this group")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should truncate long key IDs with ellipsis", () => {
|
||||
const longKeyId = "a".repeat(25);
|
||||
mockUseAccessGroupDetails.mockReturnValue({
|
||||
...baseMockReturnValue,
|
||||
data: createMockAccessGroup({ assigned_key_ids: [longKeyId] }),
|
||||
} as ReturnType<typeof useAccessGroupDetails>);
|
||||
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
expect(screen.getByText(/a{10}\.\.\.a{6}/)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display created and last updated timestamps", () => {
|
||||
renderWithProviders(
|
||||
<AccessGroupDetail accessGroupId={accessGroupId} onBack={mockOnBack} />
|
||||
);
|
||||
|
||||
expect(screen.getByText("Created")).toBeInTheDocument();
|
||||
expect(screen.getByText("Last Updated")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,345 @@
|
|||
import { useAccessGroupDetails } from "@/app/(dashboard)/hooks/accessGroups/useAccessGroupDetails";
|
||||
import {
|
||||
Button,
|
||||
Card,
|
||||
Col,
|
||||
Descriptions,
|
||||
Empty,
|
||||
Flex,
|
||||
Layout,
|
||||
List,
|
||||
Row,
|
||||
Spin,
|
||||
Tabs,
|
||||
Tag,
|
||||
theme,
|
||||
Typography
|
||||
} from "antd";
|
||||
import {
|
||||
ArrowLeftIcon,
|
||||
BotIcon,
|
||||
EditIcon,
|
||||
KeyIcon,
|
||||
LayersIcon,
|
||||
ServerIcon,
|
||||
UsersIcon,
|
||||
} from "lucide-react";
|
||||
import { useState } from "react";
|
||||
import DefaultProxyAdminTag from "../common_components/DefaultProxyAdminTag";
|
||||
import { AccessGroupEditModal } from "./AccessGroupsModal/AccessGroupEditModal";
|
||||
|
||||
const { Title, Text } = Typography;
|
||||
const { Content } = Layout;
|
||||
|
||||
interface AccessGroupDetailProps {
|
||||
accessGroupId: string;
|
||||
onBack: () => void;
|
||||
}
|
||||
|
||||
export function AccessGroupDetail({
|
||||
accessGroupId,
|
||||
onBack,
|
||||
}: AccessGroupDetailProps) {
|
||||
const { data: accessGroup, isLoading } =
|
||||
useAccessGroupDetails(accessGroupId);
|
||||
const { token } = theme.useToken();
|
||||
const [isEditModalVisible, setIsEditModalVisible] = useState(false);
|
||||
const [showAllKeys, setShowAllKeys] = useState(false);
|
||||
const [showAllTeams, setShowAllTeams] = useState(false);
|
||||
|
||||
const MAX_PREVIEW = 5;
|
||||
|
||||
if (isLoading) {
|
||||
return (
|
||||
<Content
|
||||
style={{
|
||||
padding: token.paddingLG,
|
||||
paddingInline: token.paddingLG * 2,
|
||||
}}
|
||||
>
|
||||
<Flex justify="center" align="center" style={{ minHeight: 300 }}>
|
||||
<Spin size="large" />
|
||||
</Flex>
|
||||
</Content>
|
||||
);
|
||||
}
|
||||
|
||||
if (!accessGroup) {
|
||||
return (
|
||||
<Content
|
||||
style={{
|
||||
padding: token.paddingLG,
|
||||
paddingInline: token.paddingLG * 2,
|
||||
}}
|
||||
>
|
||||
<Button
|
||||
icon={<ArrowLeftIcon size={16} />}
|
||||
onClick={onBack}
|
||||
type="text"
|
||||
style={{ marginBottom: 16 }}
|
||||
/>
|
||||
<Empty description="Access group not found" />
|
||||
</Content>
|
||||
);
|
||||
}
|
||||
|
||||
const modelIds = accessGroup.access_model_ids ?? [];
|
||||
const mcpServerIds = accessGroup.access_mcp_server_ids ?? [];
|
||||
const agentIds = accessGroup.access_agent_ids ?? [];
|
||||
const keyIds = accessGroup.assigned_key_ids ?? [];
|
||||
const teamIds = accessGroup.assigned_team_ids ?? [];
|
||||
|
||||
const displayedKeys = showAllKeys ? keyIds : keyIds.slice(0, MAX_PREVIEW);
|
||||
const displayedTeams = showAllTeams
|
||||
? teamIds
|
||||
: teamIds.slice(0, MAX_PREVIEW);
|
||||
|
||||
const handleEdit = () => {
|
||||
setIsEditModalVisible(true);
|
||||
};
|
||||
|
||||
const tabItems = [
|
||||
{
|
||||
key: "models",
|
||||
label: (
|
||||
<Flex align="center" gap={8}>
|
||||
<LayersIcon size={16} />
|
||||
Models
|
||||
<Tag style={{ marginInlineEnd: 0 }}>{modelIds.length}</Tag>
|
||||
</Flex>
|
||||
),
|
||||
children:
|
||||
modelIds.length > 0 ? (
|
||||
<List
|
||||
grid={{ gutter: 16, xs: 1, sm: 2, md: 3, lg: 4 }}
|
||||
dataSource={modelIds}
|
||||
renderItem={(id) => (
|
||||
<List.Item>
|
||||
<Card size="small">
|
||||
<Text code>{id}</Text>
|
||||
</Card>
|
||||
</List.Item>
|
||||
)}
|
||||
/>
|
||||
) : (
|
||||
<Empty description="No models assigned to this group" />
|
||||
),
|
||||
},
|
||||
{
|
||||
key: "mcp",
|
||||
label: (
|
||||
<Flex align="center" gap={8}>
|
||||
<ServerIcon size={16} />
|
||||
MCP Servers
|
||||
<Tag>{mcpServerIds.length}</Tag>
|
||||
</Flex>
|
||||
),
|
||||
children:
|
||||
mcpServerIds.length > 0 ? (
|
||||
<List
|
||||
grid={{ gutter: 16, xs: 1, sm: 2, md: 3, lg: 4 }}
|
||||
dataSource={mcpServerIds}
|
||||
renderItem={(id) => (
|
||||
<List.Item>
|
||||
<Card size="small">
|
||||
<Text code>{id}</Text>
|
||||
</Card>
|
||||
</List.Item>
|
||||
)}
|
||||
/>
|
||||
) : (
|
||||
<Empty description="No MCP servers assigned to this group" />
|
||||
),
|
||||
},
|
||||
{
|
||||
key: "agents",
|
||||
label: (
|
||||
<Flex align="center" gap={8}>
|
||||
<BotIcon size={16} />
|
||||
Agents
|
||||
<Tag>{agentIds.length}</Tag>
|
||||
</Flex>
|
||||
),
|
||||
children:
|
||||
agentIds.length > 0 ? (
|
||||
<List
|
||||
grid={{ gutter: 16, xs: 1, sm: 2, md: 3, lg: 4 }}
|
||||
dataSource={agentIds}
|
||||
renderItem={(id) => (
|
||||
<List.Item>
|
||||
<Card size="small">
|
||||
<Text code>{id}</Text>
|
||||
</Card>
|
||||
</List.Item>
|
||||
)}
|
||||
/>
|
||||
) : (
|
||||
<Empty description="No agents assigned to this group" />
|
||||
),
|
||||
},
|
||||
];
|
||||
|
||||
return (
|
||||
<Content
|
||||
style={{ padding: token.paddingLG, paddingInline: token.paddingLG * 2 }}
|
||||
>
|
||||
{/* Header */}
|
||||
<div
|
||||
style={{
|
||||
display: "flex",
|
||||
justifyContent: "space-between",
|
||||
alignItems: "center",
|
||||
marginBottom: 24,
|
||||
}}
|
||||
>
|
||||
<div style={{ display: "flex", alignItems: "center", gap: 16 }}>
|
||||
<Button
|
||||
icon={<ArrowLeftIcon size={16} />}
|
||||
onClick={onBack}
|
||||
type="text"
|
||||
/>
|
||||
<div>
|
||||
<Title level={2} style={{ margin: 0 }}>
|
||||
{accessGroup.access_group_name}
|
||||
</Title>
|
||||
<Text type="secondary">
|
||||
ID: <Text copyable>{accessGroup.access_group_id}</Text>
|
||||
</Text>
|
||||
</div>
|
||||
</div>
|
||||
<Button
|
||||
type="primary"
|
||||
icon={<EditIcon size={16} />}
|
||||
onClick={handleEdit}
|
||||
>
|
||||
Edit Access Group
|
||||
</Button>
|
||||
</div>
|
||||
|
||||
{/* Group Details */}
|
||||
<Row style={{ marginBottom: 24 }}>
|
||||
<Card>
|
||||
<Descriptions title="Group Details" column={1}>
|
||||
<Descriptions.Item label="Description">
|
||||
{accessGroup.description || "—"}
|
||||
</Descriptions.Item>
|
||||
<Descriptions.Item label="Created">
|
||||
{new Date(accessGroup.created_at).toLocaleString()}
|
||||
{accessGroup.created_by && (
|
||||
<Text>
|
||||
{"by"}
|
||||
<DefaultProxyAdminTag userId={accessGroup.created_by} />
|
||||
</Text>
|
||||
)}
|
||||
</Descriptions.Item>
|
||||
<Descriptions.Item label="Last Updated">
|
||||
{new Date(accessGroup.updated_at).toLocaleString()}
|
||||
{accessGroup.updated_by && (
|
||||
<Text>
|
||||
{"by"}
|
||||
<DefaultProxyAdminTag userId={accessGroup.updated_by} />
|
||||
</Text>
|
||||
)}
|
||||
</Descriptions.Item>
|
||||
</Descriptions>
|
||||
</Card>
|
||||
</Row>
|
||||
|
||||
{/* Attached Keys & Teams */}
|
||||
<Row gutter={[16, 16]} style={{ marginBottom: 24 }}>
|
||||
<Col xs={24} lg={12}>
|
||||
<Card
|
||||
title={
|
||||
<Flex align="center" gap={8}>
|
||||
<KeyIcon size={16} />
|
||||
Attached Keys
|
||||
<Tag>{keyIds.length}</Tag>
|
||||
</Flex>
|
||||
}
|
||||
extra={
|
||||
keyIds.length > MAX_PREVIEW ? (
|
||||
<Button
|
||||
type="link"
|
||||
onClick={() => setShowAllKeys(!showAllKeys)}
|
||||
>
|
||||
{showAllKeys ? "Show Less" : `View All (${keyIds.length})`}
|
||||
</Button>
|
||||
) : null
|
||||
}
|
||||
>
|
||||
{keyIds.length > 0 ? (
|
||||
<Flex wrap="wrap" gap={8}>
|
||||
{displayedKeys.map((id) => (
|
||||
<Tag key={id}>
|
||||
<Text code style={{ fontSize: 12 }}>
|
||||
{id.length > 20
|
||||
? `${id.slice(0, 10)}...${id.slice(-6)}`
|
||||
: id}
|
||||
</Text>
|
||||
</Tag>
|
||||
))}
|
||||
</Flex>
|
||||
) : (
|
||||
<Empty
|
||||
description="No keys attached"
|
||||
image={Empty.PRESENTED_IMAGE_SIMPLE}
|
||||
/>
|
||||
)}
|
||||
</Card>
|
||||
</Col>
|
||||
<Col xs={24} lg={12}>
|
||||
<Card
|
||||
title={
|
||||
<Flex align="center" gap={8}>
|
||||
<UsersIcon size={16} />
|
||||
Attached Teams
|
||||
<Tag>{teamIds.length}</Tag>
|
||||
</Flex>
|
||||
}
|
||||
extra={
|
||||
teamIds.length > MAX_PREVIEW ? (
|
||||
<Button
|
||||
type="link"
|
||||
onClick={() => setShowAllTeams(!showAllTeams)}
|
||||
>
|
||||
{showAllTeams
|
||||
? "Show Less"
|
||||
: `View All (${teamIds.length})`}
|
||||
</Button>
|
||||
) : null
|
||||
}
|
||||
>
|
||||
{teamIds.length > 0 ? (
|
||||
<Flex wrap="wrap" gap={8}>
|
||||
{displayedTeams.map((id) => (
|
||||
<Tag key={id}>
|
||||
<Text code style={{ fontSize: 12 }}>
|
||||
{id}
|
||||
</Text>
|
||||
</Tag>
|
||||
))}
|
||||
</Flex>
|
||||
) : (
|
||||
<Empty
|
||||
description="No teams attached"
|
||||
image={Empty.PRESENTED_IMAGE_SIMPLE}
|
||||
/>
|
||||
)}
|
||||
</Card>
|
||||
</Col>
|
||||
</Row>
|
||||
|
||||
{/* Resources Tabs */}
|
||||
<Card>
|
||||
<Tabs defaultActiveKey="models" items={tabItems} />
|
||||
</Card>
|
||||
|
||||
{/* Edit Modal */}
|
||||
<AccessGroupEditModal
|
||||
visible={isEditModalVisible}
|
||||
accessGroup={accessGroup}
|
||||
onCancel={() => setIsEditModalVisible(false)}
|
||||
/>
|
||||
</Content>
|
||||
);
|
||||
}
|
||||
|
|
@ -0,0 +1,159 @@
|
|||
import { useAgents } from "@/app/(dashboard)/hooks/agents/useAgents";
|
||||
import { useMCPServers } from "@/app/(dashboard)/hooks/mcpServers/useMCPServers";
|
||||
import { ModelSelect } from "@/components/ModelSelect/ModelSelect";
|
||||
import type { FormInstance } from "antd";
|
||||
import { Form, Input, Select, Space, Tabs } from "antd";
|
||||
import { BotIcon, InfoIcon, LayersIcon, ServerIcon } from "lucide-react";
|
||||
|
||||
const { TextArea } = Input;
|
||||
|
||||
export interface AccessGroupFormValues {
|
||||
name: string;
|
||||
description: string;
|
||||
modelIds: string[];
|
||||
mcpServerIds: string[];
|
||||
agentIds: string[];
|
||||
}
|
||||
|
||||
interface AccessGroupBaseFormProps {
|
||||
form: FormInstance<AccessGroupFormValues>;
|
||||
isNameDisabled?: boolean;
|
||||
}
|
||||
|
||||
export function AccessGroupBaseForm({
|
||||
form,
|
||||
isNameDisabled = false,
|
||||
}: AccessGroupBaseFormProps) {
|
||||
const { data: agentsData } = useAgents();
|
||||
const { data: mcpServersData } = useMCPServers();
|
||||
|
||||
const agents = agentsData?.agents ?? [];
|
||||
const mcpServers = mcpServersData ?? [];
|
||||
const items = [
|
||||
{
|
||||
key: "1",
|
||||
label: (
|
||||
<Space align="center" size={4}>
|
||||
<InfoIcon size={16} />
|
||||
General Info
|
||||
</Space>
|
||||
),
|
||||
children: (
|
||||
<div style={{ paddingTop: 16 }}>
|
||||
<Form.Item
|
||||
name="name"
|
||||
label="Group Name"
|
||||
rules={[
|
||||
{
|
||||
required: true,
|
||||
message: "Please enter the access group name",
|
||||
},
|
||||
]}
|
||||
>
|
||||
<Input
|
||||
placeholder="e.g. Engineering Team"
|
||||
disabled={isNameDisabled}
|
||||
/>
|
||||
</Form.Item>
|
||||
<Form.Item
|
||||
name="description"
|
||||
label="Description"
|
||||
>
|
||||
<TextArea
|
||||
rows={4}
|
||||
placeholder="Describe the purpose of this access group..."
|
||||
/>
|
||||
</Form.Item>
|
||||
</div>
|
||||
),
|
||||
},
|
||||
{
|
||||
key: "2",
|
||||
label: (
|
||||
<Space align="center" size={4}>
|
||||
<LayersIcon size={16} />
|
||||
Models
|
||||
</Space>
|
||||
),
|
||||
children: (
|
||||
<div style={{ paddingTop: 16 }}>
|
||||
<Form.Item name="modelIds" label="Allowed Models">
|
||||
<ModelSelect
|
||||
context="global"
|
||||
value={form.getFieldValue("modelIds") ?? []}
|
||||
onChange={(values) => form.setFieldsValue({ modelIds: values })}
|
||||
style={{ width: "100%" }}
|
||||
/>
|
||||
</Form.Item>
|
||||
</div>
|
||||
),
|
||||
},
|
||||
{
|
||||
key: "3",
|
||||
label: (
|
||||
<Space align="center" size={4}>
|
||||
<ServerIcon size={16} />
|
||||
MCP Servers
|
||||
</Space>
|
||||
),
|
||||
children: (
|
||||
<div style={{ paddingTop: 16 }}>
|
||||
<Form.Item name="mcpServerIds" label="Allowed MCP Servers">
|
||||
<Select
|
||||
mode="multiple"
|
||||
placeholder="Select MCP servers"
|
||||
style={{ width: "100%" }}
|
||||
optionFilterProp="label"
|
||||
allowClear
|
||||
options={mcpServers.map((server) => ({
|
||||
label: server.server_name ?? server.server_id,
|
||||
value: server.server_id,
|
||||
}))}
|
||||
/>
|
||||
</Form.Item>
|
||||
</div>
|
||||
),
|
||||
},
|
||||
{
|
||||
key: "4",
|
||||
label: (
|
||||
<Space align="center" size={4}>
|
||||
<BotIcon size={16} />
|
||||
Agents
|
||||
</Space>
|
||||
),
|
||||
children: (
|
||||
<div style={{ paddingTop: 16 }}>
|
||||
<Form.Item name="agentIds" label="Allowed Agents">
|
||||
<Select
|
||||
mode="multiple"
|
||||
placeholder="Select agents"
|
||||
style={{ width: "100%" }}
|
||||
optionFilterProp="label"
|
||||
allowClear
|
||||
options={agents.map((agent) => ({
|
||||
label: agent.agent_name,
|
||||
value: agent.agent_id,
|
||||
}))}
|
||||
/>
|
||||
</Form.Item>
|
||||
</div>
|
||||
),
|
||||
},
|
||||
];
|
||||
|
||||
return (
|
||||
<Form
|
||||
form={form}
|
||||
layout="vertical"
|
||||
name="access_group_form"
|
||||
initialValues={{
|
||||
modelIds: [],
|
||||
mcpServerIds: [],
|
||||
agentIds: [],
|
||||
}}
|
||||
>
|
||||
<Tabs defaultActiveKey="1" items={items} />
|
||||
</Form>
|
||||
);
|
||||
}
|
||||
|
|
@ -0,0 +1,67 @@
|
|||
import React from "react";
|
||||
import { Modal, Form, message } from "antd";
|
||||
import {
|
||||
AccessGroupBaseForm,
|
||||
AccessGroupFormValues,
|
||||
} from "./AccessGroupBaseForm";
|
||||
import {
|
||||
useCreateAccessGroup,
|
||||
AccessGroupCreateParams,
|
||||
} from "@/app/(dashboard)/hooks/accessGroups/useCreateAccessGroup";
|
||||
|
||||
interface AccessGroupCreateModalProps {
|
||||
visible: boolean;
|
||||
onCancel: () => void;
|
||||
onSuccess?: () => void;
|
||||
}
|
||||
|
||||
export function AccessGroupCreateModal({
|
||||
visible,
|
||||
onCancel,
|
||||
onSuccess,
|
||||
}: AccessGroupCreateModalProps) {
|
||||
const [form] = Form.useForm<AccessGroupFormValues>();
|
||||
const createMutation = useCreateAccessGroup();
|
||||
|
||||
const handleOk = () => {
|
||||
form
|
||||
.validateFields()
|
||||
.then((values) => {
|
||||
const params: AccessGroupCreateParams = {
|
||||
access_group_name: values.name,
|
||||
description: values.description,
|
||||
access_model_ids: values.modelIds,
|
||||
access_mcp_server_ids: values.mcpServerIds,
|
||||
access_agent_ids: values.agentIds,
|
||||
};
|
||||
|
||||
createMutation.mutate(params, {
|
||||
onSuccess: () => {
|
||||
message.success("Access group created successfully");
|
||||
form.resetFields();
|
||||
onSuccess?.();
|
||||
onCancel();
|
||||
},
|
||||
});
|
||||
})
|
||||
.catch((info) => {
|
||||
console.log("Validate Failed:", info);
|
||||
});
|
||||
};
|
||||
|
||||
return (
|
||||
<Modal
|
||||
title="Create Access Group"
|
||||
open={visible}
|
||||
onOk={handleOk}
|
||||
onCancel={onCancel}
|
||||
width={700}
|
||||
okText="Create Group"
|
||||
cancelText="Cancel"
|
||||
confirmLoading={createMutation.isPending}
|
||||
destroyOnClose
|
||||
>
|
||||
<AccessGroupBaseForm form={form} />
|
||||
</Modal>
|
||||
);
|
||||
}
|
||||
|
|
@ -0,0 +1,85 @@
|
|||
import React, { useEffect } from "react";
|
||||
import { Modal, Form, message } from "antd";
|
||||
import {
|
||||
AccessGroupBaseForm,
|
||||
AccessGroupFormValues,
|
||||
} from "./AccessGroupBaseForm";
|
||||
import {
|
||||
useEditAccessGroup,
|
||||
AccessGroupUpdateParams,
|
||||
} from "@/app/(dashboard)/hooks/accessGroups/useEditAccessGroup";
|
||||
import { AccessGroupResponse } from "@/app/(dashboard)/hooks/accessGroups/useAccessGroups";
|
||||
|
||||
interface AccessGroupEditModalProps {
|
||||
visible: boolean;
|
||||
accessGroup: AccessGroupResponse;
|
||||
onCancel: () => void;
|
||||
onSuccess?: () => void;
|
||||
}
|
||||
|
||||
export function AccessGroupEditModal({
|
||||
visible,
|
||||
accessGroup,
|
||||
onCancel,
|
||||
onSuccess,
|
||||
}: AccessGroupEditModalProps) {
|
||||
const [form] = Form.useForm<AccessGroupFormValues>();
|
||||
const editMutation = useEditAccessGroup();
|
||||
|
||||
// Populate the form with initial values whenever the modal opens or the data changes
|
||||
useEffect(() => {
|
||||
if (visible && accessGroup) {
|
||||
form.setFieldsValue({
|
||||
name: accessGroup.access_group_name,
|
||||
description: accessGroup.description ?? "",
|
||||
modelIds: accessGroup.access_model_ids ?? [],
|
||||
mcpServerIds: accessGroup.access_mcp_server_ids ?? [],
|
||||
agentIds: accessGroup.access_agent_ids ?? [],
|
||||
});
|
||||
}
|
||||
}, [visible, accessGroup, form]);
|
||||
|
||||
const handleOk = () => {
|
||||
form
|
||||
.validateFields()
|
||||
.then((values) => {
|
||||
const params: AccessGroupUpdateParams = {
|
||||
access_group_name: values.name,
|
||||
description: values.description,
|
||||
access_model_ids: values.modelIds,
|
||||
access_mcp_server_ids: values.mcpServerIds,
|
||||
access_agent_ids: values.agentIds,
|
||||
};
|
||||
|
||||
editMutation.mutate(
|
||||
{ accessGroupId: accessGroup.access_group_id, params },
|
||||
{
|
||||
onSuccess: () => {
|
||||
message.success("Access group updated successfully");
|
||||
onSuccess?.();
|
||||
onCancel();
|
||||
},
|
||||
},
|
||||
);
|
||||
})
|
||||
.catch((info) => {
|
||||
console.log("Validate Failed:", info);
|
||||
});
|
||||
};
|
||||
|
||||
return (
|
||||
<Modal
|
||||
title="Edit Access Group"
|
||||
open={visible}
|
||||
onOk={handleOk}
|
||||
onCancel={onCancel}
|
||||
width={700}
|
||||
okText="Save Changes"
|
||||
cancelText="Cancel"
|
||||
confirmLoading={editMutation.isPending}
|
||||
destroyOnHidden
|
||||
>
|
||||
<AccessGroupBaseForm form={form} />
|
||||
</Modal>
|
||||
);
|
||||
}
|
||||
|
|
@ -0,0 +1,321 @@
|
|||
import { renderWithProviders, screen, within } from "@/../tests/test-utils";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import { AccessGroupsPage } from "./AccessGroupsPage";
|
||||
import type { AccessGroupResponse } from "@/app/(dashboard)/hooks/accessGroups/useAccessGroups";
|
||||
|
||||
const mockAccessGroups: AccessGroupResponse[] = [
|
||||
{
|
||||
access_group_id: "ag-1",
|
||||
access_group_name: "Admin Group",
|
||||
description: "Administrators with full access",
|
||||
access_model_ids: ["m1", "m2"],
|
||||
access_mcp_server_ids: ["s1"],
|
||||
access_agent_ids: ["a1"],
|
||||
assigned_team_ids: [],
|
||||
assigned_key_ids: [],
|
||||
created_at: "2024-01-15T10:00:00Z",
|
||||
created_by: "user-1",
|
||||
updated_at: "2024-01-20T12:00:00Z",
|
||||
updated_by: "user-1",
|
||||
},
|
||||
{
|
||||
access_group_id: "ag-2",
|
||||
access_group_name: "Read Only",
|
||||
description: "Read-only access to models",
|
||||
access_model_ids: ["m1"],
|
||||
access_mcp_server_ids: [],
|
||||
access_agent_ids: [],
|
||||
assigned_team_ids: [],
|
||||
assigned_key_ids: [],
|
||||
created_at: "2024-01-10T09:00:00Z",
|
||||
created_by: null,
|
||||
updated_at: "2024-01-12T11:00:00Z",
|
||||
updated_by: null,
|
||||
},
|
||||
];
|
||||
|
||||
const mockUseAccessGroups = vi.fn();
|
||||
const mockUseDeleteAccessGroup = vi.fn();
|
||||
const mockMutate = vi.fn();
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/accessGroups/useAccessGroups", () => ({
|
||||
useAccessGroups: () => mockUseAccessGroups(),
|
||||
}));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/accessGroups/useDeleteAccessGroup", () => ({
|
||||
useDeleteAccessGroup: () => mockUseDeleteAccessGroup(),
|
||||
}));
|
||||
|
||||
vi.mock("./AccessGroupsDetailsPage", () => ({
|
||||
AccessGroupDetail: ({
|
||||
accessGroupId,
|
||||
onBack,
|
||||
}: {
|
||||
accessGroupId: string;
|
||||
onBack: () => void;
|
||||
}) => (
|
||||
<div data-testid="access-group-detail">
|
||||
<span>Detail for {accessGroupId}</span>
|
||||
<button onClick={onBack}>Back</button>
|
||||
</div>
|
||||
),
|
||||
}));
|
||||
|
||||
vi.mock("./AccessGroupsModal/AccessGroupCreateModal", () => ({
|
||||
AccessGroupCreateModal: ({
|
||||
visible,
|
||||
onCancel,
|
||||
}: {
|
||||
visible: boolean;
|
||||
onCancel: () => void;
|
||||
}) =>
|
||||
visible ? (
|
||||
<div data-testid="create-access-group-modal">
|
||||
<button onClick={onCancel}>Cancel</button>
|
||||
</div>
|
||||
) : null,
|
||||
}));
|
||||
|
||||
vi.mock("../common_components/IconActionButton/TableIconActionButtons/TableIconActionButton", () => ({
|
||||
default: ({
|
||||
variant,
|
||||
tooltipText,
|
||||
onClick,
|
||||
}: {
|
||||
variant: string;
|
||||
tooltipText: string;
|
||||
onClick: () => void;
|
||||
}) => (
|
||||
<button
|
||||
data-testid={`action-button-${variant.toLowerCase()}`}
|
||||
aria-label={tooltipText}
|
||||
onClick={onClick}
|
||||
>
|
||||
{variant}
|
||||
</button>
|
||||
),
|
||||
}));
|
||||
|
||||
describe("AccessGroupsPage", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
mockUseAccessGroups.mockReturnValue({
|
||||
data: mockAccessGroups,
|
||||
isLoading: false,
|
||||
});
|
||||
mockUseDeleteAccessGroup.mockReturnValue({
|
||||
mutate: mockMutate,
|
||||
isPending: false,
|
||||
});
|
||||
});
|
||||
|
||||
it("should render", () => {
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
expect(screen.getByRole("heading", { name: "Access Groups" })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display page title and subtitle", () => {
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
expect(screen.getByRole("heading", { name: "Access Groups" })).toBeInTheDocument();
|
||||
expect(
|
||||
screen.getByText("Manage resource permissions for your organization"),
|
||||
).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display Create Access Group button", () => {
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
expect(
|
||||
screen.getByRole("button", { name: /create access group/i }),
|
||||
).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display search input with placeholder", () => {
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
expect(
|
||||
screen.getByPlaceholderText("Search groups by name, ID, or description..."),
|
||||
).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display access groups in table", () => {
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
expect(screen.getByText("ag-1")).toBeInTheDocument();
|
||||
expect(screen.getByText("Admin Group")).toBeInTheDocument();
|
||||
expect(screen.getByText("ag-2")).toBeInTheDocument();
|
||||
expect(screen.getByText("Read Only")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display resource counts for each group", () => {
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
const table = screen.getByRole("table");
|
||||
expect(table).toHaveTextContent("2");
|
||||
expect(table).toHaveTextContent("1");
|
||||
});
|
||||
|
||||
it("should filter groups by search text matching name", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
const searchInput = screen.getByPlaceholderText(
|
||||
"Search groups by name, ID, or description...",
|
||||
);
|
||||
await user.type(searchInput, "Admin");
|
||||
expect(screen.getByText("Admin Group")).toBeInTheDocument();
|
||||
expect(screen.queryByText("Read Only")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should filter groups by search text matching ID", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
const searchInput = screen.getByPlaceholderText(
|
||||
"Search groups by name, ID, or description...",
|
||||
);
|
||||
await user.type(searchInput, "ag-2");
|
||||
expect(screen.getByText("Read Only")).toBeInTheDocument();
|
||||
expect(screen.queryByText("Admin Group")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should filter groups by search text matching description", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
const searchInput = screen.getByPlaceholderText(
|
||||
"Search groups by name, ID, or description...",
|
||||
);
|
||||
await user.type(searchInput, "read-only");
|
||||
expect(screen.getByText("Read Only")).toBeInTheDocument();
|
||||
expect(screen.queryByText("Admin Group")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should reset to first page when search text changes", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
const searchInput = screen.getByPlaceholderText(
|
||||
"Search groups by name, ID, or description...",
|
||||
);
|
||||
await user.type(searchInput, "Admin");
|
||||
const pagination = screen.getByText(/groups/);
|
||||
expect(pagination).toHaveTextContent("1 groups");
|
||||
});
|
||||
|
||||
it("should open create modal when Create Access Group button is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
await user.click(screen.getByRole("button", { name: /create access group/i }));
|
||||
expect(screen.getByTestId("create-access-group-modal")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should close create modal when cancel is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
await user.click(screen.getByRole("button", { name: /create access group/i }));
|
||||
expect(screen.getByTestId("create-access-group-modal")).toBeInTheDocument();
|
||||
await user.click(screen.getByRole("button", { name: "Cancel" }));
|
||||
expect(screen.queryByTestId("create-access-group-modal")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should navigate to detail view when group ID is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
await user.click(screen.getByText("ag-1"));
|
||||
expect(screen.getByTestId("access-group-detail")).toBeInTheDocument();
|
||||
expect(screen.getByText("Detail for ag-1")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should return to list view when Back is clicked from detail", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
await user.click(screen.getByText("ag-1"));
|
||||
expect(screen.getByTestId("access-group-detail")).toBeInTheDocument();
|
||||
await user.click(screen.getByRole("button", { name: "Back" }));
|
||||
expect(screen.queryByTestId("access-group-detail")).not.toBeInTheDocument();
|
||||
expect(screen.getByText("Admin Group")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should open delete modal when delete action is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
const deleteButtons = screen.getAllByRole("button", {
|
||||
name: "Delete access group",
|
||||
});
|
||||
await user.click(deleteButtons[0]);
|
||||
const dialog = screen.getByRole("dialog", { name: "Delete Access Group" });
|
||||
expect(dialog).toBeInTheDocument();
|
||||
expect(
|
||||
within(dialog).getByText(
|
||||
"Are you sure you want to delete this access group? This action cannot be undone.",
|
||||
),
|
||||
).toBeInTheDocument();
|
||||
expect(within(dialog).getByText("Access Group Information")).toBeInTheDocument();
|
||||
expect(within(dialog).getByText("ag-1")).toBeInTheDocument();
|
||||
expect(within(dialog).getByText("Admin Group")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should close delete modal when cancel is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
const deleteButtons = screen.getAllByRole("button", {
|
||||
name: "Delete access group",
|
||||
});
|
||||
await user.click(deleteButtons[0]);
|
||||
const dialog = screen.getByRole("dialog", { name: "Delete Access Group" });
|
||||
await user.click(within(dialog).getByRole("button", { name: "Cancel" }));
|
||||
expect(screen.queryByRole("dialog", { name: "Delete Access Group" })).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should call delete mutation when delete is confirmed", async () => {
|
||||
const user = userEvent.setup();
|
||||
mockMutate.mockImplementation((_id: string, opts?: { onSuccess?: () => void }) => {
|
||||
opts?.onSuccess?.();
|
||||
});
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
const deleteButtons = screen.getAllByRole("button", {
|
||||
name: "Delete access group",
|
||||
});
|
||||
await user.click(deleteButtons[0]);
|
||||
const dialog = screen.getByRole("dialog", { name: "Delete Access Group" });
|
||||
const deleteConfirmButton = within(dialog).getByRole("button", { name: /delete/i });
|
||||
await user.click(deleteConfirmButton);
|
||||
expect(mockMutate).toHaveBeenCalledWith("ag-1", expect.any(Object));
|
||||
});
|
||||
|
||||
it("should display pagination with total count", () => {
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
expect(screen.getByText("2 groups")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show table headers for ID, Name, Resources, and Actions", () => {
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
expect(screen.getByRole("columnheader", { name: /ID/i })).toBeInTheDocument();
|
||||
expect(screen.getByRole("columnheader", { name: /Name/i })).toBeInTheDocument();
|
||||
expect(screen.getByRole("columnheader", { name: /Resources/i })).toBeInTheDocument();
|
||||
expect(screen.getByRole("columnheader", { name: /Actions/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display loading state when data is loading", () => {
|
||||
mockUseAccessGroups.mockReturnValue({
|
||||
data: undefined,
|
||||
isLoading: true,
|
||||
});
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
const table = screen.getByRole("table");
|
||||
expect(table).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display empty state when no groups match search", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
const searchInput = screen.getByPlaceholderText(
|
||||
"Search groups by name, ID, or description...",
|
||||
);
|
||||
await user.type(searchInput, "nonexistent-group-xyz");
|
||||
expect(screen.getByRole("table")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display empty data when useAccessGroups returns empty array", () => {
|
||||
mockUseAccessGroups.mockReturnValue({
|
||||
data: [],
|
||||
isLoading: false,
|
||||
});
|
||||
renderWithProviders(<AccessGroupsPage />);
|
||||
expect(screen.getByRole("table")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,401 @@
|
|||
import {
|
||||
AccessGroupResponse,
|
||||
useAccessGroups,
|
||||
} from "@/app/(dashboard)/hooks/accessGroups/useAccessGroups";
|
||||
import { useDeleteAccessGroup } from "@/app/(dashboard)/hooks/accessGroups/useDeleteAccessGroup";
|
||||
import { PlusOutlined } from "@ant-design/icons";
|
||||
import {
|
||||
ColumnDef,
|
||||
flexRender,
|
||||
getCoreRowModel,
|
||||
getSortedRowModel,
|
||||
Row,
|
||||
SortingState,
|
||||
useReactTable,
|
||||
} from "@tanstack/react-table";
|
||||
import {
|
||||
Button,
|
||||
Card,
|
||||
Flex,
|
||||
Input,
|
||||
Layout,
|
||||
Pagination,
|
||||
Space,
|
||||
Table,
|
||||
Tag,
|
||||
theme,
|
||||
Tooltip,
|
||||
Typography,
|
||||
} from "antd";
|
||||
import {
|
||||
BotIcon,
|
||||
LayersIcon,
|
||||
SearchIcon,
|
||||
ServerIcon
|
||||
} from "lucide-react";
|
||||
import { useEffect, useMemo, useState } from "react";
|
||||
import DeleteResourceModal from "../common_components/DeleteResourceModal";
|
||||
import TableIconActionButton from "../common_components/IconActionButton/TableIconActionButtons/TableIconActionButton";
|
||||
import {
|
||||
SortState,
|
||||
TableHeaderSortDropdown,
|
||||
} from "../common_components/TableHeaderSortDropdown/TableHeaderSortDropdown";
|
||||
import { AccessGroupDetail } from "./AccessGroupsDetailsPage";
|
||||
import { AccessGroupCreateModal } from "./AccessGroupsModal/AccessGroupCreateModal";
|
||||
import { AccessGroup } from "./types";
|
||||
|
||||
declare module "@tanstack/react-table" {
|
||||
// eslint-disable-next-line @typescript-eslint/no-unused-vars
|
||||
interface ColumnMeta<TData, TValue> {
|
||||
responsive?: string[];
|
||||
}
|
||||
}
|
||||
|
||||
const { Title, Text } = Typography;
|
||||
const { Content } = Layout;
|
||||
|
||||
function mapResponseToAccessGroup(r: AccessGroupResponse): AccessGroup {
|
||||
return {
|
||||
id: r.access_group_id,
|
||||
name: r.access_group_name,
|
||||
description: r.description ?? "",
|
||||
modelIds: r.access_model_ids,
|
||||
mcpServerIds: r.access_mcp_server_ids,
|
||||
agentIds: r.access_agent_ids,
|
||||
keyIds: r.assigned_key_ids,
|
||||
teamIds: r.assigned_team_ids,
|
||||
createdAt: r.created_at,
|
||||
createdBy: r.created_by ?? "",
|
||||
updatedAt: r.updated_at,
|
||||
updatedBy: r.updated_by ?? "",
|
||||
};
|
||||
}
|
||||
function buildAntdColumns(
|
||||
table: ReturnType<typeof useReactTable<AccessGroup>>,
|
||||
rowLookup: Map<string, Row<AccessGroup>>,
|
||||
onSortingChange: (s: SortingState) => void,
|
||||
) {
|
||||
const headers = table.getHeaderGroups()[0]?.headers ?? [];
|
||||
|
||||
return headers.map((header) => {
|
||||
const canSort = header.column.getCanSort();
|
||||
const isSorted = header.column.getIsSorted();
|
||||
const meta = header.column.columnDef.meta as
|
||||
| { responsive?: string[] }
|
||||
| undefined;
|
||||
|
||||
const col: Record<string, unknown> = {
|
||||
title: (
|
||||
<div style={{ display: "flex", alignItems: "center", gap: 4 }}>
|
||||
{header.isPlaceholder
|
||||
? null
|
||||
: flexRender(header.column.columnDef.header, header.getContext())}
|
||||
{canSort && (
|
||||
<TableHeaderSortDropdown
|
||||
sortState={isSorted === false ? false : (isSorted as SortState)}
|
||||
onSortChange={(newState) => {
|
||||
if (newState === false) {
|
||||
onSortingChange([]);
|
||||
} else {
|
||||
onSortingChange([
|
||||
{ id: header.column.id, desc: newState === "desc" },
|
||||
]);
|
||||
}
|
||||
}}
|
||||
columnId={header.column.id}
|
||||
/>
|
||||
)}
|
||||
</div>
|
||||
),
|
||||
key: header.id,
|
||||
width: header.column.columnDef.size,
|
||||
render: (_: unknown, record: AccessGroup) => {
|
||||
const row = rowLookup.get(record.id);
|
||||
if (!row) return null;
|
||||
const cell = row
|
||||
.getVisibleCells()
|
||||
.find((c) => c.column.id === header.id);
|
||||
if (!cell) return null;
|
||||
return flexRender(cell.column.columnDef.cell, cell.getContext());
|
||||
},
|
||||
};
|
||||
|
||||
if (meta?.responsive) {
|
||||
col.responsive = meta.responsive;
|
||||
}
|
||||
|
||||
return col;
|
||||
});
|
||||
}
|
||||
|
||||
export function AccessGroupsPage() {
|
||||
const { token } = theme.useToken();
|
||||
const { data: groupsData, isLoading } = useAccessGroups();
|
||||
const groups = useMemo(
|
||||
() => (groupsData ?? []).map(mapResponseToAccessGroup),
|
||||
[groupsData],
|
||||
);
|
||||
|
||||
const [selectedGroupId, setSelectedGroupId] = useState<string | null>(null);
|
||||
const [isCreateModalVisible, setIsCreateModalVisible] = useState(false);
|
||||
const [searchText, setSearchText] = useState("");
|
||||
const [currentPage, setCurrentPage] = useState(1);
|
||||
const [sorting, setSorting] = useState<SortingState>([]);
|
||||
const [groupToDelete, setGroupToDelete] = useState<AccessGroup | null>(null);
|
||||
const deleteMutation = useDeleteAccessGroup();
|
||||
const pageSize = 10;
|
||||
|
||||
useEffect(() => {
|
||||
setCurrentPage(1);
|
||||
}, [searchText]);
|
||||
|
||||
// ---------- filtered data ----------
|
||||
const filteredGroups = useMemo(
|
||||
() =>
|
||||
groups.filter(
|
||||
(group) =>
|
||||
group.name.toLowerCase().includes(searchText.toLowerCase()) ||
|
||||
group.id.toLowerCase().includes(searchText.toLowerCase()) ||
|
||||
group.description.toLowerCase().includes(searchText.toLowerCase()),
|
||||
),
|
||||
[groups, searchText],
|
||||
);
|
||||
|
||||
// ---------- TanStack column definitions ----------
|
||||
const columnDefs = useMemo<ColumnDef<AccessGroup>[]>(
|
||||
() => [
|
||||
{
|
||||
id: "id",
|
||||
accessorKey: "id",
|
||||
header: () => <span>ID</span>,
|
||||
enableSorting: false,
|
||||
size: 170,
|
||||
cell: ({ row }) => {
|
||||
const record = row.original;
|
||||
return (
|
||||
<Tooltip title={record.id}>
|
||||
<Text
|
||||
ellipsis
|
||||
className="text-blue-500 bg-blue-50 hover:bg-blue-100 text-xs cursor-pointer"
|
||||
style={{ fontSize: 14, padding: "1px 8px" }}
|
||||
onClick={() => setSelectedGroupId(record.id)}
|
||||
>
|
||||
{record.id}
|
||||
</Text>
|
||||
</Tooltip>
|
||||
);
|
||||
},
|
||||
},
|
||||
{
|
||||
id: "name",
|
||||
accessorKey: "name",
|
||||
header: () => <span>Name</span>,
|
||||
enableSorting: true,
|
||||
cell: ({ getValue }) => getValue() as string,
|
||||
},
|
||||
{
|
||||
id: "resources",
|
||||
header: () => <span>Resources</span>,
|
||||
enableSorting: false,
|
||||
cell: ({ row }) => {
|
||||
const record = row.original;
|
||||
return (
|
||||
<Flex gap={12} align="center">
|
||||
<Tooltip title={`${record.modelIds.length} Models`}>
|
||||
<Tag color="blue" style={{ fontSize: 14, padding: "2px 8px", margin: 0 }}>
|
||||
<Flex align="center" gap={6}>
|
||||
<LayersIcon size={14} />
|
||||
{record.modelIds.length}
|
||||
</Flex>
|
||||
</Tag>
|
||||
</Tooltip>
|
||||
<Tooltip title={`${record.mcpServerIds.length} MCP Servers`}>
|
||||
<Tag color="cyan" style={{ fontSize: 14, padding: "2px 8px", margin: 0 }}>
|
||||
<Flex align="center" gap={6}>
|
||||
<ServerIcon size={14} />
|
||||
{record.mcpServerIds.length}
|
||||
</Flex>
|
||||
</Tag>
|
||||
</Tooltip>
|
||||
<Tooltip title={`${record.agentIds.length} Agents`}>
|
||||
<Tag color="purple" style={{ fontSize: 14, padding: "2px 8px", margin: 0 }}>
|
||||
<Flex align="center" gap={6}>
|
||||
<BotIcon size={14} />
|
||||
{record.agentIds.length}
|
||||
</Flex>
|
||||
</Tag>
|
||||
</Tooltip>
|
||||
</Flex>
|
||||
);
|
||||
},
|
||||
},
|
||||
{
|
||||
id: "createdAt",
|
||||
accessorKey: "createdAt",
|
||||
header: () => <span>Created</span>,
|
||||
enableSorting: true,
|
||||
sortingFn: "datetime",
|
||||
cell: ({ getValue }) =>
|
||||
new Date(getValue() as string).toLocaleDateString(),
|
||||
meta: { responsive: ["lg"] },
|
||||
},
|
||||
{
|
||||
id: "updatedAt",
|
||||
accessorKey: "updatedAt",
|
||||
header: () => <span>Updated</span>,
|
||||
enableSorting: false,
|
||||
cell: ({ getValue }) =>
|
||||
new Date(getValue() as string).toLocaleDateString(),
|
||||
meta: { responsive: ["xl"] },
|
||||
},
|
||||
{
|
||||
id: "actions",
|
||||
header: () => <span>Actions</span>,
|
||||
enableSorting: false,
|
||||
cell: ({ row }) => (
|
||||
<Space>
|
||||
<TableIconActionButton
|
||||
variant="Delete"
|
||||
tooltipText="Delete access group"
|
||||
onClick={() => setGroupToDelete(row.original)}
|
||||
/>
|
||||
</Space>
|
||||
),
|
||||
},
|
||||
],
|
||||
// setSelectedGroup is stable (useState setter)
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
[],
|
||||
);
|
||||
|
||||
// ---------- TanStack table instance ----------
|
||||
const table = useReactTable<AccessGroup>({
|
||||
data: filteredGroups,
|
||||
columns: columnDefs,
|
||||
state: { sorting },
|
||||
onSortingChange: setSorting,
|
||||
getCoreRowModel: getCoreRowModel(),
|
||||
getSortedRowModel: getSortedRowModel(),
|
||||
getRowId: (row) => row.id,
|
||||
});
|
||||
|
||||
// All sorted rows from TanStack
|
||||
const sortedRows = table.getRowModel().rows;
|
||||
|
||||
// Paginated slice
|
||||
const paginatedRows = sortedRows.slice(
|
||||
(currentPage - 1) * pageSize,
|
||||
currentPage * pageSize,
|
||||
);
|
||||
|
||||
// Map for O(1) lookup by record id in antd render()
|
||||
const rowLookup = useMemo(
|
||||
() => new Map(paginatedRows.map((row) => [row.original.id, row])),
|
||||
[paginatedRows],
|
||||
);
|
||||
|
||||
// Convert TanStack headers → antd columns
|
||||
const antdColumns = buildAntdColumns(table, rowLookup, setSorting);
|
||||
|
||||
// antd dataSource (just the originals for the current page)
|
||||
const dataSource = paginatedRows.map((row) => row.original);
|
||||
|
||||
if (selectedGroupId) {
|
||||
return (
|
||||
<AccessGroupDetail
|
||||
accessGroupId={selectedGroupId}
|
||||
onBack={() => setSelectedGroupId(null)}
|
||||
/>
|
||||
);
|
||||
}
|
||||
|
||||
return (
|
||||
<Content
|
||||
style={{ padding: token.paddingLG, paddingInline: token.paddingLG * 2 }}
|
||||
>
|
||||
<Flex
|
||||
justify="space-between"
|
||||
align="center"
|
||||
style={{ marginBottom: 16 }}
|
||||
>
|
||||
<Space direction="vertical" size={0}>
|
||||
<Title level={2} style={{ margin: 0 }}>
|
||||
Access Groups
|
||||
</Title>
|
||||
<Text type="secondary">
|
||||
Manage resource permissions for your organization
|
||||
</Text>
|
||||
</Space>
|
||||
<Button
|
||||
type="primary"
|
||||
icon={<PlusOutlined />}
|
||||
onClick={() => setIsCreateModalVisible(true)}
|
||||
>
|
||||
Create Access Group
|
||||
</Button>
|
||||
</Flex>
|
||||
|
||||
<Card styles={{ body: { padding: 0 } }}>
|
||||
<Flex
|
||||
justify="space-between"
|
||||
align="center"
|
||||
style={{
|
||||
padding: "12px 16px",
|
||||
}}
|
||||
>
|
||||
<Input
|
||||
prefix={<SearchIcon size={16} />}
|
||||
placeholder="Search groups by name, ID, or description..."
|
||||
style={{ maxWidth: 400 }}
|
||||
value={searchText}
|
||||
onChange={(e) => setSearchText(e.target.value)}
|
||||
allowClear
|
||||
/>
|
||||
<Pagination
|
||||
current={currentPage}
|
||||
total={sortedRows.length}
|
||||
pageSize={pageSize}
|
||||
onChange={(page) => setCurrentPage(page)}
|
||||
size="small"
|
||||
showTotal={(total) => `${total} groups`}
|
||||
showSizeChanger={false}
|
||||
/>
|
||||
</Flex>
|
||||
<Table
|
||||
columns={antdColumns}
|
||||
dataSource={dataSource}
|
||||
rowKey="id"
|
||||
loading={isLoading}
|
||||
pagination={false}
|
||||
/>
|
||||
</Card>
|
||||
|
||||
<AccessGroupCreateModal
|
||||
visible={isCreateModalVisible}
|
||||
onCancel={() => setIsCreateModalVisible(false)}
|
||||
/>
|
||||
|
||||
<DeleteResourceModal
|
||||
isOpen={!!groupToDelete}
|
||||
title="Delete Access Group"
|
||||
message="Are you sure you want to delete this access group? This action cannot be undone."
|
||||
resourceInformationTitle="Access Group Information"
|
||||
resourceInformation={[
|
||||
{ label: "ID", value: groupToDelete?.id, code: true },
|
||||
{ label: "Name", value: groupToDelete?.name },
|
||||
{ label: "Description", value: groupToDelete?.description || "—" },
|
||||
]}
|
||||
onCancel={() => setGroupToDelete(null)}
|
||||
onOk={() => {
|
||||
if (!groupToDelete) return;
|
||||
deleteMutation.mutate(groupToDelete.id, {
|
||||
onSuccess: () => {
|
||||
setGroupToDelete(null);
|
||||
},
|
||||
});
|
||||
}}
|
||||
confirmLoading={deleteMutation.isPending}
|
||||
/>
|
||||
</Content>
|
||||
);
|
||||
}
|
||||
46
ui/litellm-dashboard/src/components/AccessGroups/types.ts
Normal file
46
ui/litellm-dashboard/src/components/AccessGroups/types.ts
Normal file
|
|
@ -0,0 +1,46 @@
|
|||
export interface AccessGroup {
|
||||
id: string
|
||||
name: string
|
||||
description: string
|
||||
modelIds: string[]
|
||||
mcpServerIds: string[]
|
||||
agentIds: string[]
|
||||
keyIds: string[]
|
||||
teamIds: string[]
|
||||
createdAt: string
|
||||
createdBy: string
|
||||
updatedAt: string
|
||||
updatedBy: string
|
||||
}
|
||||
|
||||
export interface Model {
|
||||
id: string
|
||||
name: string
|
||||
provider: string
|
||||
}
|
||||
|
||||
export interface McpServer {
|
||||
id: string
|
||||
name: string
|
||||
endpoint: string
|
||||
}
|
||||
|
||||
export interface Agent {
|
||||
id: string
|
||||
name: string
|
||||
type: string
|
||||
}
|
||||
|
||||
export interface AccessGroupKey {
|
||||
id: string
|
||||
alias: string
|
||||
status: string
|
||||
createdAt: string
|
||||
}
|
||||
|
||||
export interface AccessGroupTeam {
|
||||
id: string
|
||||
name: string
|
||||
members: number
|
||||
role: string
|
||||
}
|
||||
|
|
@ -0,0 +1,24 @@
|
|||
import { Tag, Typography } from "antd";
|
||||
|
||||
const { Text } = Typography;
|
||||
|
||||
const DEFAULT_USER_ID = "default_user_id";
|
||||
|
||||
interface DefaultProxyAdminTagProps {
|
||||
userId: string | null | undefined;
|
||||
}
|
||||
|
||||
/**
|
||||
* Renders "Default Proxy Admin" as a blue Tag when the given userId is
|
||||
* the well-known `default_user_id`, otherwise renders the raw value as
|
||||
* plain text.
|
||||
*/
|
||||
export default function DefaultProxyAdminTag({
|
||||
userId,
|
||||
}: DefaultProxyAdminTagProps) {
|
||||
if (userId === DEFAULT_USER_ID) {
|
||||
return <Tag color="blue">Default Proxy Admin</Tag>;
|
||||
}
|
||||
|
||||
return <Text>{userId}</Text>;
|
||||
}
|
||||
|
|
@ -151,11 +151,7 @@ const menuGroups: MenuGroup[] = [
|
|||
{
|
||||
key: "logs",
|
||||
page: "logs",
|
||||
label: (
|
||||
<span className="flex items-center gap-4">
|
||||
Logs <NewBadge />
|
||||
</span>
|
||||
),
|
||||
label: "Logs",
|
||||
icon: <LineChartOutlined />,
|
||||
},
|
||||
],
|
||||
|
|
@ -183,6 +179,17 @@ const menuGroups: MenuGroup[] = [
|
|||
icon: <BankOutlined />,
|
||||
roles: all_admin_roles,
|
||||
},
|
||||
{
|
||||
key: "access-groups",
|
||||
page: "access-groups",
|
||||
label: (
|
||||
<span className="flex items-center gap-2">
|
||||
Access Groups <NewBadge />
|
||||
</span>
|
||||
),
|
||||
icon: <BlockOutlined />,
|
||||
roles: all_admin_roles,
|
||||
},
|
||||
{
|
||||
key: "budgets",
|
||||
page: "budgets",
|
||||
|
|
|
|||
|
|
@ -5673,6 +5673,37 @@ export const deletePolicyAttachmentCall = async (accessToken: string, attachment
|
|||
}
|
||||
};
|
||||
|
||||
export const testPipelineCall = async (
|
||||
accessToken: string,
|
||||
pipeline: any,
|
||||
testMessages: Array<{role: string, content: string}>
|
||||
) => {
|
||||
try {
|
||||
const url = proxyBaseUrl ? `${proxyBaseUrl}/policies/test-pipeline` : `/policies/test-pipeline`;
|
||||
const response = await fetch(url, {
|
||||
method: "POST",
|
||||
headers: {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
body: JSON.stringify({ pipeline, test_messages: testMessages }),
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorData = await response.json();
|
||||
const errorMessage = deriveErrorMessage(errorData);
|
||||
handleError(errorMessage);
|
||||
throw new Error(errorMessage);
|
||||
}
|
||||
|
||||
const data = await response.json();
|
||||
return data;
|
||||
} catch (error) {
|
||||
console.error("Failed to test pipeline:", error);
|
||||
throw error;
|
||||
}
|
||||
};
|
||||
|
||||
export const getResolvedGuardrails = async (accessToken: string, policyId: string) => {
|
||||
try {
|
||||
const url = proxyBaseUrl
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ export const pageDescriptions: Record<string, string> = {
|
|||
users: "Manage internal user accounts and permissions",
|
||||
teams: "Create and manage teams for access control",
|
||||
organizations: "Manage organizations and their members",
|
||||
"access-groups": "Manage access groups for role-based permissions",
|
||||
budgets: "Set and monitor spending budgets",
|
||||
api_ref: "Browse API documentation and endpoints",
|
||||
"model-hub-table": "Explore available AI models and providers",
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ interface AddPolicyFormProps {
|
|||
visible: boolean;
|
||||
onClose: () => void;
|
||||
onSuccess: () => void;
|
||||
onOpenFlowBuilder: () => void;
|
||||
accessToken: string | null;
|
||||
editingPolicy?: Policy | null;
|
||||
existingPolicies: Policy[];
|
||||
|
|
@ -22,10 +23,117 @@ interface AddPolicyFormProps {
|
|||
updatePolicy: (accessToken: string, policyId: string, policyData: any) => Promise<any>;
|
||||
}
|
||||
|
||||
// ─────────────────────────────────────────────────────────────────────────────
|
||||
// Mode Picker (Step 1) - shown first when creating a new policy
|
||||
// ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
interface ModePicker {
|
||||
selected: "simple" | "flow_builder";
|
||||
onSelect: (mode: "simple" | "flow_builder") => void;
|
||||
}
|
||||
|
||||
const ModePicker: React.FC<ModePicker> = ({ selected, onSelect }) => (
|
||||
<div className="flex gap-4" style={{ padding: "8px 0" }}>
|
||||
{/* Simple Mode Card */}
|
||||
<div
|
||||
onClick={() => onSelect("simple")}
|
||||
style={{
|
||||
flex: 1,
|
||||
padding: "24px 20px",
|
||||
border: `2px solid ${selected === "simple" ? "#4f46e5" : "#e5e7eb"}`,
|
||||
borderRadius: 12,
|
||||
cursor: "pointer",
|
||||
backgroundColor: selected === "simple" ? "#eef2ff" : "#fff",
|
||||
transition: "all 0.15s ease",
|
||||
}}
|
||||
>
|
||||
<div
|
||||
style={{
|
||||
width: 40,
|
||||
height: 40,
|
||||
borderRadius: 10,
|
||||
backgroundColor: selected === "simple" ? "#e0e7ff" : "#f3f4f6",
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
justifyContent: "center",
|
||||
marginBottom: 16,
|
||||
}}
|
||||
>
|
||||
<svg width="20" height="20" viewBox="0 0 24 24" fill="none" stroke={selected === "simple" ? "#4f46e5" : "#6b7280"} strokeWidth="2" strokeLinecap="round" strokeLinejoin="round">
|
||||
<rect x="3" y="3" width="18" height="18" rx="2" />
|
||||
<path d="M8 7h8M8 12h8M8 17h5" />
|
||||
</svg>
|
||||
</div>
|
||||
<Text strong style={{ fontSize: 15, display: "block", marginBottom: 4 }}>
|
||||
Simple Mode
|
||||
</Text>
|
||||
<Text type="secondary" style={{ fontSize: 13 }}>
|
||||
Pick guardrails from a list. All run in parallel.
|
||||
</Text>
|
||||
</div>
|
||||
|
||||
{/* Flow Builder Card */}
|
||||
<div
|
||||
onClick={() => onSelect("flow_builder")}
|
||||
style={{
|
||||
flex: 1,
|
||||
padding: "24px 20px",
|
||||
border: `2px solid ${selected === "flow_builder" ? "#4f46e5" : "#e5e7eb"}`,
|
||||
borderRadius: 12,
|
||||
cursor: "pointer",
|
||||
backgroundColor: selected === "flow_builder" ? "#eef2ff" : "#fff",
|
||||
transition: "all 0.15s ease",
|
||||
position: "relative",
|
||||
}}
|
||||
>
|
||||
<Tag
|
||||
color="purple"
|
||||
style={{
|
||||
position: "absolute",
|
||||
top: 12,
|
||||
right: 12,
|
||||
fontSize: 10,
|
||||
fontWeight: 600,
|
||||
margin: 0,
|
||||
}}
|
||||
>
|
||||
NEW
|
||||
</Tag>
|
||||
<div
|
||||
style={{
|
||||
width: 40,
|
||||
height: 40,
|
||||
borderRadius: 10,
|
||||
backgroundColor: selected === "flow_builder" ? "#e0e7ff" : "#f3f4f6",
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
justifyContent: "center",
|
||||
marginBottom: 16,
|
||||
}}
|
||||
>
|
||||
<svg width="20" height="20" viewBox="0 0 24 24" fill="none" stroke={selected === "flow_builder" ? "#4f46e5" : "#6b7280"} strokeWidth="2" strokeLinecap="round" strokeLinejoin="round">
|
||||
<path d="M13 2L3 14h9l-1 8 10-12h-9l1-8z" />
|
||||
</svg>
|
||||
</div>
|
||||
<Text strong style={{ fontSize: 15, display: "block", marginBottom: 4 }}>
|
||||
Flow Builder
|
||||
</Text>
|
||||
<Text type="secondary" style={{ fontSize: 13 }}>
|
||||
Define steps, conditions, and error responses.
|
||||
</Text>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
|
||||
// ─────────────────────────────────────────────────────────────────────────────
|
||||
// Main Component
|
||||
// ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
const AddPolicyForm: React.FC<AddPolicyFormProps> = ({
|
||||
visible,
|
||||
onClose,
|
||||
onSuccess,
|
||||
onOpenFlowBuilder,
|
||||
accessToken,
|
||||
editingPolicy,
|
||||
existingPolicies,
|
||||
|
|
@ -39,16 +147,16 @@ const AddPolicyForm: React.FC<AddPolicyFormProps> = ({
|
|||
const [isLoadingResolved, setIsLoadingResolved] = useState(false);
|
||||
const [modelConditionType, setModelConditionType] = useState<"model" | "regex">("model");
|
||||
const [availableModels, setAvailableModels] = useState<string[]>([]);
|
||||
const [step, setStep] = useState<"pick_mode" | "simple_form">("pick_mode");
|
||||
const [selectedMode, setSelectedMode] = useState<"simple" | "flow_builder">("simple");
|
||||
const { userId, userRole } = useAuthorized();
|
||||
|
||||
// Only consider it "editing" if editingPolicy has a policy_id (real existing policy)
|
||||
// If editingPolicy is set but has no policy_id, it's just pre-filled data for a new policy (e.g., from a template)
|
||||
const isEditing = !!editingPolicy?.policy_id;
|
||||
|
||||
useEffect(() => {
|
||||
if (visible && editingPolicy) {
|
||||
const modelCondition = editingPolicy.condition?.model;
|
||||
// Detect if it's a regex pattern (contains *, ., [, ], etc.)
|
||||
const isRegex = modelCondition && /[.*+?^${}()|[\]\\]/.test(modelCondition);
|
||||
setModelConditionType(isRegex ? "regex" : "model");
|
||||
|
||||
|
|
@ -60,14 +168,25 @@ const AddPolicyForm: React.FC<AddPolicyFormProps> = ({
|
|||
guardrails_remove: editingPolicy.guardrails_remove || [],
|
||||
model_condition: modelCondition,
|
||||
});
|
||||
// Load resolved guardrails for editing
|
||||
|
||||
if (editingPolicy.policy_id && accessToken) {
|
||||
loadResolvedGuardrails(editingPolicy.policy_id);
|
||||
}
|
||||
|
||||
// If editing a pipeline policy, go directly to flow builder
|
||||
if (editingPolicy.pipeline) {
|
||||
onClose();
|
||||
onOpenFlowBuilder();
|
||||
return;
|
||||
}
|
||||
// If editing a simple policy, skip mode picker
|
||||
setStep("simple_form");
|
||||
} else if (visible) {
|
||||
form.resetFields();
|
||||
setResolvedGuardrails([]);
|
||||
setModelConditionType("model");
|
||||
setSelectedMode("simple");
|
||||
setStep("pick_mode");
|
||||
}
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [visible, editingPolicy, form]);
|
||||
|
|
@ -81,7 +200,6 @@ const AddPolicyForm: React.FC<AddPolicyFormProps> = ({
|
|||
|
||||
const loadAvailableModels = async () => {
|
||||
if (!accessToken) return;
|
||||
|
||||
try {
|
||||
const response = await modelAvailableCall(accessToken, userId, userRole);
|
||||
if (response?.data) {
|
||||
|
|
@ -95,7 +213,6 @@ const AddPolicyForm: React.FC<AddPolicyFormProps> = ({
|
|||
|
||||
const loadResolvedGuardrails = async (policyId: string) => {
|
||||
if (!accessToken) return;
|
||||
|
||||
setIsLoadingResolved(true);
|
||||
try {
|
||||
const data = await getResolvedGuardrails(accessToken, policyId);
|
||||
|
|
@ -115,20 +232,15 @@ const AddPolicyForm: React.FC<AddPolicyFormProps> = ({
|
|||
|
||||
let resolved = new Set<string>();
|
||||
|
||||
// If inheriting, find parent policy and get its guardrails
|
||||
if (inheritFrom) {
|
||||
const parentPolicy = existingPolicies.find(p => p.policy_name === inheritFrom);
|
||||
if (parentPolicy) {
|
||||
// Recursively resolve parent's guardrails
|
||||
const parentResolved = resolveParentGuardrails(parentPolicy);
|
||||
parentResolved.forEach(g => resolved.add(g));
|
||||
}
|
||||
}
|
||||
|
||||
// Add guardrails
|
||||
guardrailsAdd.forEach((g: string) => resolved.add(g));
|
||||
|
||||
// Remove guardrails
|
||||
guardrailsRemove.forEach((g: string) => resolved.delete(g));
|
||||
|
||||
return Array.from(resolved).sort();
|
||||
|
|
@ -137,32 +249,23 @@ const AddPolicyForm: React.FC<AddPolicyFormProps> = ({
|
|||
const resolveParentGuardrails = (policy: Policy): string[] => {
|
||||
let resolved = new Set<string>();
|
||||
|
||||
// If parent inherits, resolve recursively
|
||||
if (policy.inherit) {
|
||||
const grandparent = existingPolicies.find(p => p.policy_name === policy.inherit);
|
||||
if (grandparent) {
|
||||
const grandparentResolved = resolveParentGuardrails(grandparent);
|
||||
grandparentResolved.forEach(g => resolved.add(g));
|
||||
resolveParentGuardrails(grandparent).forEach(g => resolved.add(g));
|
||||
}
|
||||
}
|
||||
|
||||
// Add parent's guardrails
|
||||
if (policy.guardrails_add) {
|
||||
policy.guardrails_add.forEach(g => resolved.add(g));
|
||||
}
|
||||
|
||||
// Remove parent's removed guardrails
|
||||
if (policy.guardrails_remove) {
|
||||
policy.guardrails_remove.forEach(g => resolved.delete(g));
|
||||
}
|
||||
|
||||
return Array.from(resolved);
|
||||
};
|
||||
|
||||
// Recompute resolved guardrails when form values change
|
||||
const handleFormChange = () => {
|
||||
const resolved = computeResolvedGuardrails();
|
||||
setResolvedGuardrails(resolved);
|
||||
setResolvedGuardrails(computeResolvedGuardrails());
|
||||
};
|
||||
|
||||
const resetForm = () => {
|
||||
|
|
@ -171,9 +274,20 @@ const AddPolicyForm: React.FC<AddPolicyFormProps> = ({
|
|||
|
||||
const handleClose = () => {
|
||||
resetForm();
|
||||
setStep("pick_mode");
|
||||
setSelectedMode("simple");
|
||||
onClose();
|
||||
};
|
||||
|
||||
const handleModeConfirm = () => {
|
||||
if (selectedMode === "flow_builder") {
|
||||
onClose();
|
||||
onOpenFlowBuilder();
|
||||
} else {
|
||||
setStep("simple_form");
|
||||
}
|
||||
};
|
||||
|
||||
const handleSubmit = async () => {
|
||||
try {
|
||||
setIsSubmitting(true);
|
||||
|
|
@ -228,6 +342,50 @@ const AddPolicyForm: React.FC<AddPolicyFormProps> = ({
|
|||
value: p.policy_name,
|
||||
}));
|
||||
|
||||
// ── Mode Picker Step ──────────────────────────────────────────────────────
|
||||
if (step === "pick_mode") {
|
||||
return (
|
||||
<Modal
|
||||
title="Create New Policy"
|
||||
open={visible}
|
||||
onCancel={handleClose}
|
||||
footer={null}
|
||||
width={620}
|
||||
>
|
||||
<ModePicker selected={selectedMode} onSelect={setSelectedMode} />
|
||||
|
||||
{selectedMode === "flow_builder" && (
|
||||
<Alert
|
||||
message="You'll be redirected to the full-screen Flow Builder to design your policy logic visually."
|
||||
type="info"
|
||||
style={{
|
||||
marginTop: 16,
|
||||
backgroundColor: "#eef2ff",
|
||||
border: "1px solid #c7d2fe",
|
||||
}}
|
||||
/>
|
||||
)}
|
||||
|
||||
<div className="flex justify-end gap-2" style={{ marginTop: 24 }}>
|
||||
<Button variant="secondary" onClick={handleClose}>
|
||||
Cancel
|
||||
</Button>
|
||||
<Button
|
||||
onClick={handleModeConfirm}
|
||||
style={{
|
||||
backgroundColor: "#4f46e5",
|
||||
color: "#fff",
|
||||
border: "none",
|
||||
}}
|
||||
>
|
||||
{selectedMode === "flow_builder" ? "Continue to Builder" : "Create Policy"}
|
||||
</Button>
|
||||
</div>
|
||||
</Modal>
|
||||
);
|
||||
}
|
||||
|
||||
// ── Simple Form Step ──────────────────────────────────────────────────────
|
||||
return (
|
||||
<Modal
|
||||
title={isEditing ? "Edit Policy" : "Create New Policy"}
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import { isAdminRole } from "@/utils/roles";
|
|||
import PolicyTable from "./policy_table";
|
||||
import PolicyInfoView from "./policy_info";
|
||||
import AddPolicyForm from "./add_policy_form";
|
||||
import { FlowBuilderPage } from "./pipeline_flow_builder";
|
||||
import AttachmentTable from "./attachment_table";
|
||||
import AddAttachmentForm from "./add_attachment_form";
|
||||
import PolicyTestPanel from "./policy_test_panel";
|
||||
|
|
@ -56,6 +57,7 @@ const PoliciesPanel: React.FC<PoliciesPanelProps> = ({
|
|||
const [selectedTemplate, setSelectedTemplate] = useState<any>(null);
|
||||
const [existingGuardrailNames, setExistingGuardrailNames] = useState<Set<string>>(new Set());
|
||||
const [isCreatingGuardrails, setIsCreatingGuardrails] = useState(false);
|
||||
const [showFlowBuilder, setShowFlowBuilder] = useState(false);
|
||||
|
||||
const isAdmin = userRole ? isAdminRole(userRole) : false;
|
||||
|
||||
|
|
@ -349,8 +351,12 @@ const PoliciesPanel: React.FC<PoliciesPanelProps> = ({
|
|||
onClose={() => setSelectedPolicyId(null)}
|
||||
onEdit={(policy) => {
|
||||
setEditingPolicy(policy);
|
||||
setIsAddPolicyModalVisible(true);
|
||||
setSelectedPolicyId(null);
|
||||
if (policy.pipeline) {
|
||||
setShowFlowBuilder(true);
|
||||
} else {
|
||||
setIsAddPolicyModalVisible(true);
|
||||
}
|
||||
}}
|
||||
accessToken={accessToken}
|
||||
isAdmin={isAdmin}
|
||||
|
|
@ -363,7 +369,11 @@ const PoliciesPanel: React.FC<PoliciesPanelProps> = ({
|
|||
onDeleteClick={handleDeleteClick}
|
||||
onEditClick={(policy) => {
|
||||
setEditingPolicy(policy);
|
||||
setIsAddPolicyModalVisible(true);
|
||||
if (policy.pipeline) {
|
||||
setShowFlowBuilder(true);
|
||||
} else {
|
||||
setIsAddPolicyModalVisible(true);
|
||||
}
|
||||
}}
|
||||
onViewClick={(policyId) => setSelectedPolicyId(policyId)}
|
||||
isAdmin={isAdmin}
|
||||
|
|
@ -374,6 +384,10 @@ const PoliciesPanel: React.FC<PoliciesPanelProps> = ({
|
|||
visible={isAddPolicyModalVisible}
|
||||
onClose={handleCloseModal}
|
||||
onSuccess={handleSuccess}
|
||||
onOpenFlowBuilder={() => {
|
||||
setIsAddPolicyModalVisible(false);
|
||||
setShowFlowBuilder(true);
|
||||
}}
|
||||
accessToken={accessToken}
|
||||
editingPolicy={editingPolicy}
|
||||
existingPolicies={policiesList}
|
||||
|
|
@ -473,6 +487,24 @@ const PoliciesPanel: React.FC<PoliciesPanelProps> = ({
|
|||
</TabPanel>
|
||||
</TabPanels>
|
||||
</TabGroup>
|
||||
|
||||
{showFlowBuilder && (
|
||||
<FlowBuilderPage
|
||||
onBack={() => {
|
||||
setShowFlowBuilder(false);
|
||||
setEditingPolicy(null);
|
||||
}}
|
||||
onSuccess={() => {
|
||||
fetchPolicies();
|
||||
setEditingPolicy(null);
|
||||
}}
|
||||
accessToken={accessToken}
|
||||
editingPolicy={editingPolicy}
|
||||
availableGuardrails={guardrailsList}
|
||||
createPolicy={createPolicyCall}
|
||||
updatePolicy={updatePolicyCall}
|
||||
/>
|
||||
)}
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -0,0 +1,999 @@
|
|||
import React, { useState } from "react";
|
||||
import { Select, Typography, message } from "antd";
|
||||
import { Button, TextInput } from "@tremor/react";
|
||||
import { ArrowLeftIcon, PlusIcon } from "@heroicons/react/outline";
|
||||
import { DotsVerticalIcon } from "@heroicons/react/solid";
|
||||
import { GuardrailPipeline, PipelineStep, PipelineTestResult, PolicyCreateRequest, PolicyUpdateRequest, Policy } from "./types";
|
||||
import { Guardrail } from "../guardrails/types";
|
||||
import { testPipelineCall } from "../networking";
|
||||
import NotificationsManager from "../molecules/notifications_manager";
|
||||
|
||||
const { Text } = Typography;
|
||||
|
||||
const ACTION_OPTIONS = [
|
||||
{ label: "Next Step", value: "next" },
|
||||
{ label: "Allow", value: "allow" },
|
||||
{ label: "Block", value: "block" },
|
||||
{ label: "Custom Response", value: "modify_response" },
|
||||
];
|
||||
|
||||
const ACTION_LABELS: Record<string, string> = {
|
||||
allow: "Allow",
|
||||
block: "Block",
|
||||
next: "Next Step",
|
||||
modify_response: "Custom Response",
|
||||
};
|
||||
|
||||
function createDefaultStep(): PipelineStep {
|
||||
return {
|
||||
guardrail: "",
|
||||
on_pass: "next",
|
||||
on_fail: "block",
|
||||
pass_data: false,
|
||||
modify_response_message: null,
|
||||
};
|
||||
}
|
||||
|
||||
function insertStep(steps: PipelineStep[], atIndex: number): PipelineStep[] {
|
||||
const newSteps = [...steps];
|
||||
newSteps.splice(atIndex, 0, createDefaultStep());
|
||||
return newSteps;
|
||||
}
|
||||
|
||||
function removeStep(steps: PipelineStep[], index: number): PipelineStep[] {
|
||||
if (steps.length <= 1) return steps;
|
||||
const newSteps = [...steps];
|
||||
newSteps.splice(index, 1);
|
||||
return newSteps;
|
||||
}
|
||||
|
||||
function updateStepAtIndex(
|
||||
steps: PipelineStep[],
|
||||
index: number,
|
||||
updated: Partial<PipelineStep>
|
||||
): PipelineStep[] {
|
||||
return steps.map((s, i) => (i === index ? { ...s, ...updated } : s));
|
||||
}
|
||||
|
||||
// ─────────────────────────────────────────────────────────────────────────────
|
||||
// Icons (matching the reference image)
|
||||
// ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
const GuardrailIcon: React.FC = () => (
|
||||
<div
|
||||
style={{
|
||||
width: 28,
|
||||
height: 28,
|
||||
borderRadius: "50%",
|
||||
backgroundColor: "#eef2ff",
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
justifyContent: "center",
|
||||
flexShrink: 0,
|
||||
}}
|
||||
>
|
||||
<svg width="14" height="14" viewBox="0 0 24 24" fill="none" stroke="#6366f1" strokeWidth="2" strokeLinecap="round" strokeLinejoin="round">
|
||||
<circle cx="12" cy="12" r="10" />
|
||||
<path d="M12 8v4" />
|
||||
</svg>
|
||||
</div>
|
||||
);
|
||||
|
||||
const PlayIcon: React.FC = () => (
|
||||
<div
|
||||
style={{
|
||||
width: 28,
|
||||
height: 28,
|
||||
borderRadius: "50%",
|
||||
backgroundColor: "#f3f4f6",
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
justifyContent: "center",
|
||||
flexShrink: 0,
|
||||
}}
|
||||
>
|
||||
<svg width="12" height="12" viewBox="0 0 24 24" fill="#6b7280" stroke="none">
|
||||
<polygon points="6,3 20,12 6,21" />
|
||||
</svg>
|
||||
</div>
|
||||
);
|
||||
|
||||
const PassIcon: React.FC = () => (
|
||||
<svg width="14" height="14" viewBox="0 0 24 24" fill="none" stroke="#22c55e" strokeWidth="2.5" strokeLinecap="round" strokeLinejoin="round" style={{ flexShrink: 0 }}>
|
||||
<circle cx="12" cy="12" r="10" />
|
||||
<path d="M9 12l2 2 4-4" />
|
||||
</svg>
|
||||
);
|
||||
|
||||
const FailIcon: React.FC = () => (
|
||||
<svg width="14" height="14" viewBox="0 0 24 24" fill="none" stroke="#f87171" strokeWidth="2.5" strokeLinecap="round" strokeLinejoin="round" style={{ flexShrink: 0 }}>
|
||||
<circle cx="12" cy="12" r="10" />
|
||||
</svg>
|
||||
);
|
||||
|
||||
// ─────────────────────────────────────────────────────────────────────────────
|
||||
// Connector
|
||||
// ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
interface ConnectorProps {
|
||||
onInsert: () => void;
|
||||
}
|
||||
|
||||
const Connector: React.FC<ConnectorProps> = ({ onInsert }) => (
|
||||
<div className="flex flex-col items-center" style={{ height: 56 }}>
|
||||
<div style={{ width: 1, flex: 1, backgroundColor: "#d1d5db" }} />
|
||||
<button
|
||||
onClick={onInsert}
|
||||
className="flex items-center justify-center"
|
||||
style={{
|
||||
width: 24,
|
||||
height: 24,
|
||||
borderRadius: "50%",
|
||||
border: "1px solid #d1d5db",
|
||||
backgroundColor: "#fff",
|
||||
cursor: "pointer",
|
||||
zIndex: 1,
|
||||
transition: "all 0.15s ease",
|
||||
}}
|
||||
onMouseEnter={(e) => {
|
||||
e.currentTarget.style.borderColor = "#6366f1";
|
||||
e.currentTarget.style.backgroundColor = "#eef2ff";
|
||||
}}
|
||||
onMouseLeave={(e) => {
|
||||
e.currentTarget.style.borderColor = "#d1d5db";
|
||||
e.currentTarget.style.backgroundColor = "#fff";
|
||||
}}
|
||||
title="Insert step"
|
||||
>
|
||||
<PlusIcon style={{ width: 12, height: 12, color: "#9ca3af" }} />
|
||||
</button>
|
||||
<div style={{ width: 1, flex: 1, backgroundColor: "#d1d5db" }} />
|
||||
</div>
|
||||
);
|
||||
|
||||
// ─────────────────────────────────────────────────────────────────────────────
|
||||
// Step Card (editable)
|
||||
// ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
interface StepCardProps {
|
||||
step: PipelineStep;
|
||||
stepIndex: number;
|
||||
totalSteps: number;
|
||||
onChange: (updated: Partial<PipelineStep>) => void;
|
||||
onDelete: () => void;
|
||||
availableGuardrails: Guardrail[];
|
||||
}
|
||||
|
||||
const StepCard: React.FC<StepCardProps> = ({
|
||||
step,
|
||||
stepIndex,
|
||||
totalSteps,
|
||||
onChange,
|
||||
onDelete,
|
||||
availableGuardrails,
|
||||
}) => {
|
||||
const guardrailOptions = availableGuardrails.map((g) => ({
|
||||
label: g.guardrail_name || g.guardrail_id,
|
||||
value: g.guardrail_name || g.guardrail_id,
|
||||
}));
|
||||
|
||||
return (
|
||||
<div
|
||||
style={{
|
||||
border: "1px solid #e5e7eb",
|
||||
borderRadius: 10,
|
||||
backgroundColor: "#fff",
|
||||
maxWidth: 720,
|
||||
width: "100%",
|
||||
overflow: "hidden",
|
||||
}}
|
||||
>
|
||||
{/* Header row */}
|
||||
<div
|
||||
className="flex items-center justify-between"
|
||||
style={{ padding: "14px 20px 0 20px" }}
|
||||
>
|
||||
<div className="flex items-center gap-2">
|
||||
<GuardrailIcon />
|
||||
<span
|
||||
style={{
|
||||
fontSize: 11,
|
||||
fontWeight: 700,
|
||||
textTransform: "uppercase",
|
||||
color: "#6366f1",
|
||||
letterSpacing: "0.06em",
|
||||
}}
|
||||
>
|
||||
GUARDRAIL
|
||||
</span>
|
||||
</div>
|
||||
<div className="flex items-center gap-2">
|
||||
<span style={{ fontSize: 13, color: "#9ca3af" }}>
|
||||
Step {stepIndex + 1}
|
||||
</span>
|
||||
<button
|
||||
onClick={onDelete}
|
||||
disabled={totalSteps <= 1}
|
||||
style={{
|
||||
background: "none",
|
||||
border: "none",
|
||||
cursor: totalSteps <= 1 ? "not-allowed" : "pointer",
|
||||
opacity: totalSteps <= 1 ? 0.3 : 1,
|
||||
padding: 2,
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
}}
|
||||
title="Delete step"
|
||||
>
|
||||
<DotsVerticalIcon style={{ width: 16, height: 16, color: "#9ca3af" }} />
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* Guardrail selector */}
|
||||
<div style={{ padding: "12px 20px 16px 20px" }}>
|
||||
<label style={{ fontSize: 12, fontWeight: 500, color: "#6b7280", display: "block", marginBottom: 6 }}>
|
||||
Guardrail
|
||||
</label>
|
||||
<Select
|
||||
showSearch
|
||||
style={{ width: "100%" }}
|
||||
placeholder="Select a guardrail"
|
||||
value={step.guardrail || undefined}
|
||||
onChange={(value) => onChange({ guardrail: value })}
|
||||
options={guardrailOptions}
|
||||
filterOption={(input, option) =>
|
||||
(option?.label ?? "").toString().toLowerCase().includes(input.toLowerCase())
|
||||
}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{/* ON PASS section */}
|
||||
<div style={{ borderTop: "1px solid #f0f0f0", padding: "14px 20px" }}>
|
||||
<div className="flex items-center gap-2" style={{ marginBottom: 8 }}>
|
||||
<PassIcon />
|
||||
<span style={{ fontSize: 13, fontWeight: 600, color: "#374151" }}>ON PASS</span>
|
||||
</div>
|
||||
<label style={{ fontSize: 12, fontWeight: 500, color: "#6b7280", display: "block", marginBottom: 6 }}>
|
||||
Action
|
||||
</label>
|
||||
<Select
|
||||
style={{ width: "100%" }}
|
||||
value={step.on_pass}
|
||||
onChange={(value) => onChange({ on_pass: value as PipelineStep["on_pass"] })}
|
||||
options={ACTION_OPTIONS}
|
||||
/>
|
||||
{step.on_pass === "modify_response" && (
|
||||
<div style={{ marginTop: 8 }}>
|
||||
<label style={{ fontSize: 12, fontWeight: 500, color: "#6b7280", display: "block", marginBottom: 6 }}>
|
||||
Custom Response Message
|
||||
</label>
|
||||
<TextInput
|
||||
placeholder="Enter custom response..."
|
||||
value={step.modify_response_message || ""}
|
||||
onChange={(e) => onChange({ modify_response_message: e.target.value || null })}
|
||||
/>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* ON FAIL section */}
|
||||
<div style={{ borderTop: "1px solid #f0f0f0", padding: "14px 20px" }}>
|
||||
<div className="flex items-center gap-2" style={{ marginBottom: 8 }}>
|
||||
<FailIcon />
|
||||
<span style={{ fontSize: 13, fontWeight: 600, color: "#374151" }}>ON FAIL</span>
|
||||
</div>
|
||||
<label style={{ fontSize: 12, fontWeight: 500, color: "#6b7280", display: "block", marginBottom: 6 }}>
|
||||
Action
|
||||
</label>
|
||||
<Select
|
||||
style={{ width: "100%" }}
|
||||
value={step.on_fail}
|
||||
onChange={(value) => onChange({ on_fail: value as PipelineStep["on_fail"] })}
|
||||
options={ACTION_OPTIONS}
|
||||
/>
|
||||
{step.on_fail === "modify_response" && (
|
||||
<div style={{ marginTop: 8 }}>
|
||||
<label style={{ fontSize: 12, fontWeight: 500, color: "#6b7280", display: "block", marginBottom: 6 }}>
|
||||
Custom Response Message
|
||||
</label>
|
||||
<TextInput
|
||||
placeholder="Enter custom response..."
|
||||
value={step.modify_response_message || ""}
|
||||
onChange={(e) => onChange({ modify_response_message: e.target.value || null })}
|
||||
/>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
// ─────────────────────────────────────────────────────────────────────────────
|
||||
// Main Component
|
||||
// ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
interface PipelineFlowBuilderProps {
|
||||
pipeline: GuardrailPipeline;
|
||||
onChange: (pipeline: GuardrailPipeline) => void;
|
||||
availableGuardrails: Guardrail[];
|
||||
}
|
||||
|
||||
const PipelineFlowBuilder: React.FC<PipelineFlowBuilderProps> = ({
|
||||
pipeline,
|
||||
onChange,
|
||||
availableGuardrails,
|
||||
}) => {
|
||||
const handleInsertStep = (atIndex: number) => {
|
||||
onChange({ ...pipeline, steps: insertStep(pipeline.steps, atIndex) });
|
||||
};
|
||||
|
||||
const handleRemoveStep = (index: number) => {
|
||||
onChange({ ...pipeline, steps: removeStep(pipeline.steps, index) });
|
||||
};
|
||||
|
||||
const handleUpdateStep = (index: number, updated: Partial<PipelineStep>) => {
|
||||
onChange({
|
||||
...pipeline,
|
||||
steps: updateStepAtIndex(pipeline.steps, index, updated),
|
||||
});
|
||||
};
|
||||
|
||||
return (
|
||||
<div className="flex flex-col items-center" style={{ padding: "16px 0" }}>
|
||||
{/* Trigger Card */}
|
||||
<div
|
||||
style={{
|
||||
border: "1px solid #e5e7eb",
|
||||
borderRadius: 10,
|
||||
padding: "16px 20px",
|
||||
backgroundColor: "#fff",
|
||||
maxWidth: 720,
|
||||
width: "100%",
|
||||
}}
|
||||
>
|
||||
<div className="flex items-center gap-3">
|
||||
<PlayIcon />
|
||||
<div>
|
||||
<span
|
||||
style={{
|
||||
fontSize: 11,
|
||||
fontWeight: 700,
|
||||
textTransform: "uppercase",
|
||||
color: "#6b7280",
|
||||
letterSpacing: "0.06em",
|
||||
display: "block",
|
||||
marginBottom: 2,
|
||||
}}
|
||||
>
|
||||
TRIGGER
|
||||
</span>
|
||||
<span style={{ fontSize: 14, fontWeight: 600, color: "#111827", display: "block" }}>
|
||||
Incoming LLM Request
|
||||
</span>
|
||||
<span style={{ fontSize: 13, color: "#9ca3af" }}>
|
||||
This flow runs when a request matches this policy
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* Steps */}
|
||||
{pipeline.steps.map((step, index) => (
|
||||
<React.Fragment key={index}>
|
||||
<Connector onInsert={() => handleInsertStep(index)} />
|
||||
<StepCard
|
||||
step={step}
|
||||
stepIndex={index}
|
||||
totalSteps={pipeline.steps.length}
|
||||
onChange={(updated) => handleUpdateStep(index, updated)}
|
||||
onDelete={() => handleRemoveStep(index)}
|
||||
availableGuardrails={availableGuardrails}
|
||||
/>
|
||||
</React.Fragment>
|
||||
))}
|
||||
|
||||
{/* Bottom connector */}
|
||||
<Connector onInsert={() => handleInsertStep(pipeline.steps.length)} />
|
||||
|
||||
{/* End card */}
|
||||
<div
|
||||
style={{
|
||||
border: "1px solid #e5e7eb",
|
||||
borderRadius: 10,
|
||||
padding: "14px 20px",
|
||||
backgroundColor: "#fff",
|
||||
maxWidth: 720,
|
||||
width: "100%",
|
||||
}}
|
||||
>
|
||||
<div className="flex items-center gap-3">
|
||||
<div
|
||||
style={{
|
||||
width: 28,
|
||||
height: 28,
|
||||
borderRadius: "50%",
|
||||
backgroundColor: "#f3f4f6",
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
justifyContent: "center",
|
||||
flexShrink: 0,
|
||||
}}
|
||||
>
|
||||
<svg width="12" height="12" viewBox="0 0 24 24" fill="none" stroke="#6b7280" strokeWidth="2.5" strokeLinecap="round" strokeLinejoin="round">
|
||||
<rect x="3" y="3" width="18" height="18" rx="2" />
|
||||
<line x1="8" y1="12" x2="16" y2="12" />
|
||||
</svg>
|
||||
</div>
|
||||
<div>
|
||||
<span
|
||||
style={{
|
||||
fontSize: 11,
|
||||
fontWeight: 700,
|
||||
textTransform: "uppercase",
|
||||
color: "#6b7280",
|
||||
letterSpacing: "0.06em",
|
||||
display: "block",
|
||||
marginBottom: 2,
|
||||
}}
|
||||
>
|
||||
END
|
||||
</span>
|
||||
<span style={{ fontSize: 14, fontWeight: 600, color: "#111827", display: "block" }}>
|
||||
Continue to LLM
|
||||
</span>
|
||||
<span style={{ fontSize: 13, color: "#9ca3af" }}>
|
||||
Request proceeds to the model
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
// ─────────────────────────────────────────────────────────────────────────────
|
||||
// Read-only display for policy info view
|
||||
// ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
interface PipelineInfoDisplayProps {
|
||||
pipeline: GuardrailPipeline;
|
||||
}
|
||||
|
||||
export const PipelineInfoDisplay: React.FC<PipelineInfoDisplayProps> = ({ pipeline }) => (
|
||||
<div className="flex flex-col items-center" style={{ padding: "16px 0" }}>
|
||||
{/* Trigger */}
|
||||
<div
|
||||
style={{
|
||||
border: "1px solid #e5e7eb",
|
||||
borderRadius: 10,
|
||||
padding: "14px 20px",
|
||||
backgroundColor: "#fff",
|
||||
maxWidth: 720,
|
||||
width: "100%",
|
||||
}}
|
||||
>
|
||||
<div className="flex items-center gap-3">
|
||||
<PlayIcon />
|
||||
<div>
|
||||
<span style={{ fontSize: 11, fontWeight: 700, textTransform: "uppercase", color: "#6b7280", letterSpacing: "0.06em", display: "block", marginBottom: 2 }}>
|
||||
TRIGGER
|
||||
</span>
|
||||
<span style={{ fontSize: 14, fontWeight: 600, color: "#111827" }}>
|
||||
Incoming LLM Request
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* Steps */}
|
||||
{pipeline.steps.map((step, index) => (
|
||||
<React.Fragment key={index}>
|
||||
{/* Connector */}
|
||||
<div style={{ width: 1, height: 32, backgroundColor: "#d1d5db" }} />
|
||||
|
||||
{/* Step card */}
|
||||
<div
|
||||
style={{
|
||||
border: "1px solid #e5e7eb",
|
||||
borderRadius: 10,
|
||||
padding: "14px 20px",
|
||||
backgroundColor: "#fff",
|
||||
maxWidth: 720,
|
||||
width: "100%",
|
||||
}}
|
||||
>
|
||||
{/* Header */}
|
||||
<div className="flex items-center justify-between" style={{ marginBottom: 8 }}>
|
||||
<div className="flex items-center gap-2">
|
||||
<GuardrailIcon />
|
||||
<span style={{ fontSize: 11, fontWeight: 700, textTransform: "uppercase", color: "#6366f1", letterSpacing: "0.06em" }}>
|
||||
GUARDRAIL
|
||||
</span>
|
||||
</div>
|
||||
<span style={{ fontSize: 13, color: "#9ca3af" }}>Step {index + 1}</span>
|
||||
</div>
|
||||
|
||||
{/* Name */}
|
||||
<div style={{ fontSize: 15, fontWeight: 600, color: "#111827", marginBottom: 8 }}>
|
||||
{step.guardrail}
|
||||
</div>
|
||||
|
||||
{/* Divider */}
|
||||
<div style={{ borderTop: "1px solid #f3f4f6", marginBottom: 10 }} />
|
||||
|
||||
{/* Pass / Fail */}
|
||||
<div className="flex items-center gap-6" style={{ fontSize: 13, color: "#374151" }}>
|
||||
<span className="flex items-center gap-1.5">
|
||||
<PassIcon /> Pass → {ACTION_LABELS[step.on_pass] || step.on_pass}
|
||||
</span>
|
||||
<span className="flex items-center gap-1.5">
|
||||
<FailIcon /> Fail → {ACTION_LABELS[step.on_fail] || step.on_fail}
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
</React.Fragment>
|
||||
))}
|
||||
</div>
|
||||
);
|
||||
|
||||
// ─────────────────────────────────────────────────────────────────────────────
|
||||
// Pipeline Test Panel (right drawer)
|
||||
// ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
interface PipelineTestPanelProps {
|
||||
pipeline: GuardrailPipeline;
|
||||
accessToken: string | null;
|
||||
onClose: () => void;
|
||||
}
|
||||
|
||||
const OUTCOME_STYLES: Record<string, { bg: string; color: string; label: string }> = {
|
||||
pass: { bg: "#f0fdf4", color: "#16a34a", label: "PASS" },
|
||||
fail: { bg: "#fef2f2", color: "#dc2626", label: "FAIL" },
|
||||
error: { bg: "#fffbeb", color: "#d97706", label: "ERROR" },
|
||||
};
|
||||
|
||||
const TERMINAL_STYLES: Record<string, { bg: string; color: string }> = {
|
||||
allow: { bg: "#f0fdf4", color: "#16a34a" },
|
||||
block: { bg: "#fef2f2", color: "#dc2626" },
|
||||
modify_response: { bg: "#eff6ff", color: "#2563eb" },
|
||||
};
|
||||
|
||||
const PipelineTestPanel: React.FC<PipelineTestPanelProps> = ({
|
||||
pipeline,
|
||||
accessToken,
|
||||
onClose,
|
||||
}) => {
|
||||
const [testMessage, setTestMessage] = useState("Hello, can you help me?");
|
||||
const [isRunning, setIsRunning] = useState(false);
|
||||
const [result, setResult] = useState<PipelineTestResult | null>(null);
|
||||
const [error, setError] = useState<string | null>(null);
|
||||
|
||||
const handleRunTest = async () => {
|
||||
if (!accessToken) return;
|
||||
|
||||
const emptySteps = pipeline.steps.filter((s) => !s.guardrail);
|
||||
if (emptySteps.length > 0) {
|
||||
setError("All steps must have a guardrail selected");
|
||||
return;
|
||||
}
|
||||
|
||||
setIsRunning(true);
|
||||
setResult(null);
|
||||
setError(null);
|
||||
|
||||
try {
|
||||
const data = await testPipelineCall(
|
||||
accessToken,
|
||||
pipeline,
|
||||
[{ role: "user", content: testMessage }]
|
||||
);
|
||||
setResult(data);
|
||||
} catch (e) {
|
||||
setError(e instanceof Error ? e.message : String(e));
|
||||
} finally {
|
||||
setIsRunning(false);
|
||||
}
|
||||
};
|
||||
|
||||
return (
|
||||
<div
|
||||
style={{
|
||||
width: 400,
|
||||
borderLeft: "1px solid #e5e7eb",
|
||||
backgroundColor: "#fff",
|
||||
display: "flex",
|
||||
flexDirection: "column",
|
||||
flexShrink: 0,
|
||||
overflow: "hidden",
|
||||
}}
|
||||
>
|
||||
{/* Panel header */}
|
||||
<div
|
||||
style={{
|
||||
padding: "12px 16px",
|
||||
borderBottom: "1px solid #e5e7eb",
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
justifyContent: "space-between",
|
||||
}}
|
||||
>
|
||||
<span style={{ fontSize: 14, fontWeight: 600, color: "#111827" }}>Test Pipeline</span>
|
||||
<button
|
||||
onClick={onClose}
|
||||
style={{
|
||||
background: "none",
|
||||
border: "none",
|
||||
cursor: "pointer",
|
||||
fontSize: 18,
|
||||
color: "#9ca3af",
|
||||
padding: "0 4px",
|
||||
}}
|
||||
>
|
||||
x
|
||||
</button>
|
||||
</div>
|
||||
|
||||
{/* Input section */}
|
||||
<div style={{ padding: 16, borderBottom: "1px solid #e5e7eb" }}>
|
||||
<label style={{ fontSize: 12, fontWeight: 500, color: "#6b7280", display: "block", marginBottom: 6 }}>
|
||||
Test Message
|
||||
</label>
|
||||
<textarea
|
||||
value={testMessage}
|
||||
onChange={(e) => setTestMessage(e.target.value)}
|
||||
placeholder="Enter a test message..."
|
||||
rows={3}
|
||||
style={{
|
||||
width: "100%",
|
||||
border: "1px solid #d1d5db",
|
||||
borderRadius: 6,
|
||||
padding: "8px 10px",
|
||||
fontSize: 13,
|
||||
resize: "vertical",
|
||||
fontFamily: "inherit",
|
||||
}}
|
||||
/>
|
||||
<Button
|
||||
onClick={handleRunTest}
|
||||
loading={isRunning}
|
||||
style={{ marginTop: 8, width: "100%" }}
|
||||
>
|
||||
Run Test
|
||||
</Button>
|
||||
</div>
|
||||
|
||||
{/* Results section */}
|
||||
<div style={{ flex: 1, overflowY: "auto", padding: 16 }}>
|
||||
{error && (
|
||||
<div
|
||||
style={{
|
||||
padding: "10px 12px",
|
||||
backgroundColor: "#fef2f2",
|
||||
border: "1px solid #fecaca",
|
||||
borderRadius: 6,
|
||||
fontSize: 13,
|
||||
color: "#dc2626",
|
||||
marginBottom: 12,
|
||||
}}
|
||||
>
|
||||
{error}
|
||||
</div>
|
||||
)}
|
||||
|
||||
{result && (
|
||||
<div>
|
||||
{/* Step results */}
|
||||
{result.step_results.map((step, i) => {
|
||||
const style = OUTCOME_STYLES[step.outcome] || OUTCOME_STYLES.error;
|
||||
return (
|
||||
<div
|
||||
key={i}
|
||||
style={{
|
||||
border: "1px solid #e5e7eb",
|
||||
borderRadius: 8,
|
||||
padding: "10px 12px",
|
||||
marginBottom: 8,
|
||||
}}
|
||||
>
|
||||
<div className="flex items-center justify-between" style={{ marginBottom: 4 }}>
|
||||
<span style={{ fontSize: 13, fontWeight: 600, color: "#111827" }}>
|
||||
Step {i + 1}: {step.guardrail_name}
|
||||
</span>
|
||||
<span
|
||||
style={{
|
||||
fontSize: 11,
|
||||
fontWeight: 700,
|
||||
backgroundColor: style.bg,
|
||||
color: style.color,
|
||||
padding: "2px 8px",
|
||||
borderRadius: 4,
|
||||
}}
|
||||
>
|
||||
{style.label}
|
||||
</span>
|
||||
</div>
|
||||
<div style={{ fontSize: 12, color: "#6b7280" }}>
|
||||
Action: {ACTION_LABELS[step.action_taken] || step.action_taken}
|
||||
{step.duration_seconds != null && (
|
||||
<span style={{ marginLeft: 8 }}>
|
||||
({(step.duration_seconds * 1000).toFixed(0)}ms)
|
||||
</span>
|
||||
)}
|
||||
</div>
|
||||
{step.error_detail && (
|
||||
<div style={{ fontSize: 12, color: "#dc2626", marginTop: 4 }}>
|
||||
{step.error_detail}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
);
|
||||
})}
|
||||
|
||||
{/* Terminal result */}
|
||||
<div
|
||||
style={{
|
||||
borderTop: "1px solid #e5e7eb",
|
||||
paddingTop: 12,
|
||||
marginTop: 4,
|
||||
}}
|
||||
>
|
||||
<div className="flex items-center justify-between">
|
||||
<span style={{ fontSize: 13, fontWeight: 600, color: "#111827" }}>Result</span>
|
||||
{(() => {
|
||||
const ts = TERMINAL_STYLES[result.terminal_action] || TERMINAL_STYLES.block;
|
||||
return (
|
||||
<span
|
||||
style={{
|
||||
fontSize: 12,
|
||||
fontWeight: 700,
|
||||
backgroundColor: ts.bg,
|
||||
color: ts.color,
|
||||
padding: "3px 10px",
|
||||
borderRadius: 4,
|
||||
textTransform: "uppercase",
|
||||
}}
|
||||
>
|
||||
{result.terminal_action === "modify_response" ? "Custom Response" : result.terminal_action}
|
||||
</span>
|
||||
);
|
||||
})()}
|
||||
</div>
|
||||
{result.error_message && (
|
||||
<div style={{ fontSize: 12, color: "#dc2626", marginTop: 6 }}>
|
||||
{result.error_message}
|
||||
</div>
|
||||
)}
|
||||
{result.modify_response_message && (
|
||||
<div style={{ fontSize: 12, color: "#2563eb", marginTop: 6 }}>
|
||||
Response: {result.modify_response_message}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{!result && !error && (
|
||||
<div style={{ textAlign: "center", color: "#9ca3af", fontSize: 13, marginTop: 24 }}>
|
||||
Enter a test message and click "Run Test" to execute the pipeline
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
// ─────────────────────────────────────────────────────────────────────────────
|
||||
// Full-screen Flow Builder Page
|
||||
// ─────────────────────────────────────────────────────────────────────────────
|
||||
|
||||
interface FlowBuilderPageProps {
|
||||
onBack: () => void;
|
||||
onSuccess: () => void;
|
||||
accessToken: string | null;
|
||||
editingPolicy?: Policy | null;
|
||||
availableGuardrails: Guardrail[];
|
||||
createPolicy: (accessToken: string, policyData: any) => Promise<any>;
|
||||
updatePolicy: (accessToken: string, policyId: string, policyData: any) => Promise<any>;
|
||||
}
|
||||
|
||||
export const FlowBuilderPage: React.FC<FlowBuilderPageProps> = ({
|
||||
onBack,
|
||||
onSuccess,
|
||||
accessToken,
|
||||
editingPolicy,
|
||||
availableGuardrails,
|
||||
createPolicy,
|
||||
updatePolicy,
|
||||
}) => {
|
||||
const isEditing = !!editingPolicy?.policy_id;
|
||||
|
||||
const [policyName, setPolicyName] = useState(editingPolicy?.policy_name || "");
|
||||
const [description, setDescription] = useState(editingPolicy?.description || "");
|
||||
const [isSubmitting, setIsSubmitting] = useState(false);
|
||||
const [showTestPanel, setShowTestPanel] = useState(false);
|
||||
const [pipeline, setPipeline] = useState<GuardrailPipeline>(
|
||||
editingPolicy?.pipeline || { mode: "pre_call", steps: [createDefaultStep()] }
|
||||
);
|
||||
|
||||
const handleSave = async () => {
|
||||
if (!policyName.trim()) {
|
||||
message.error("Please enter a policy name");
|
||||
return;
|
||||
}
|
||||
if (!accessToken) {
|
||||
message.error("No access token available");
|
||||
return;
|
||||
}
|
||||
|
||||
const emptySteps = pipeline.steps.filter((s) => !s.guardrail);
|
||||
if (emptySteps.length > 0) {
|
||||
message.error("Please select a guardrail for all steps");
|
||||
return;
|
||||
}
|
||||
|
||||
setIsSubmitting(true);
|
||||
try {
|
||||
const guardrailsFromPipeline = pipeline.steps
|
||||
.map((s) => s.guardrail)
|
||||
.filter(Boolean);
|
||||
|
||||
const data: PolicyCreateRequest | PolicyUpdateRequest = {
|
||||
policy_name: policyName,
|
||||
description: description || undefined,
|
||||
guardrails_add: guardrailsFromPipeline,
|
||||
guardrails_remove: [],
|
||||
pipeline: pipeline,
|
||||
};
|
||||
|
||||
if (isEditing && editingPolicy) {
|
||||
await updatePolicy(accessToken, editingPolicy.policy_id, data as PolicyUpdateRequest);
|
||||
NotificationsManager.success("Policy updated successfully");
|
||||
} else {
|
||||
await createPolicy(accessToken, data as PolicyCreateRequest);
|
||||
NotificationsManager.success("Policy created successfully");
|
||||
}
|
||||
|
||||
onSuccess();
|
||||
onBack();
|
||||
} catch (error) {
|
||||
console.error("Failed to save policy:", error);
|
||||
NotificationsManager.fromBackend(
|
||||
"Failed to save policy: " + (error instanceof Error ? error.message : String(error))
|
||||
);
|
||||
} finally {
|
||||
setIsSubmitting(false);
|
||||
}
|
||||
};
|
||||
|
||||
return (
|
||||
<div
|
||||
style={{
|
||||
position: "fixed",
|
||||
top: 0,
|
||||
left: 0,
|
||||
right: 0,
|
||||
bottom: 0,
|
||||
backgroundColor: "#f9fafb",
|
||||
zIndex: 1000,
|
||||
display: "flex",
|
||||
flexDirection: "column",
|
||||
overflow: "hidden",
|
||||
}}
|
||||
>
|
||||
{/* Header bar */}
|
||||
<div
|
||||
style={{
|
||||
borderBottom: "1px solid #e5e7eb",
|
||||
backgroundColor: "#fff",
|
||||
padding: "10px 24px",
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
justifyContent: "space-between",
|
||||
flexShrink: 0,
|
||||
}}
|
||||
>
|
||||
<div className="flex items-center gap-3">
|
||||
<button
|
||||
onClick={onBack}
|
||||
style={{
|
||||
background: "none",
|
||||
border: "none",
|
||||
cursor: "pointer",
|
||||
padding: 4,
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
}}
|
||||
>
|
||||
<ArrowLeftIcon style={{ width: 18, height: 18, color: "#6b7280" }} />
|
||||
</button>
|
||||
<span style={{ fontSize: 14, color: "#6b7280" }}>Policies</span>
|
||||
<span style={{ fontSize: 14, color: "#d1d5db" }}>/</span>
|
||||
<TextInput
|
||||
placeholder="Policy name..."
|
||||
value={policyName}
|
||||
onChange={(e) => setPolicyName(e.target.value)}
|
||||
disabled={isEditing}
|
||||
style={{ width: 240 }}
|
||||
/>
|
||||
<span
|
||||
style={{
|
||||
fontSize: 11,
|
||||
fontWeight: 600,
|
||||
backgroundColor: "#eef2ff",
|
||||
color: "#6366f1",
|
||||
padding: "3px 8px",
|
||||
borderRadius: 4,
|
||||
letterSpacing: "0.02em",
|
||||
}}
|
||||
>
|
||||
Flow
|
||||
</span>
|
||||
</div>
|
||||
<div className="flex items-center gap-2">
|
||||
<Button variant="secondary" onClick={onBack}>
|
||||
Cancel
|
||||
</Button>
|
||||
<Button
|
||||
variant="secondary"
|
||||
onClick={() => setShowTestPanel(!showTestPanel)}
|
||||
>
|
||||
{showTestPanel ? "Hide Test" : "Test Pipeline"}
|
||||
</Button>
|
||||
<Button onClick={handleSave} loading={isSubmitting}>
|
||||
{isEditing ? "Update Policy" : "Save Policy"}
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* Description bar */}
|
||||
<div
|
||||
style={{
|
||||
padding: "8px 24px",
|
||||
backgroundColor: "#fff",
|
||||
borderBottom: "1px solid #e5e7eb",
|
||||
flexShrink: 0,
|
||||
}}
|
||||
>
|
||||
<TextInput
|
||||
placeholder="Add a description (optional)..."
|
||||
value={description}
|
||||
onChange={(e) => setDescription(e.target.value)}
|
||||
style={{ maxWidth: 500 }}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{/* Flow builder canvas + test panel */}
|
||||
<div style={{ flex: 1, display: "flex", overflow: "hidden" }}>
|
||||
<div
|
||||
style={{
|
||||
flex: 1,
|
||||
overflowY: "auto",
|
||||
display: "flex",
|
||||
justifyContent: "center",
|
||||
padding: "32px 24px",
|
||||
}}
|
||||
>
|
||||
<div style={{ maxWidth: 760, width: "100%" }}>
|
||||
<PipelineFlowBuilder
|
||||
pipeline={pipeline}
|
||||
onChange={setPipeline}
|
||||
availableGuardrails={availableGuardrails}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{showTestPanel && (
|
||||
<PipelineTestPanel
|
||||
pipeline={pipeline}
|
||||
accessToken={accessToken}
|
||||
onClose={() => setShowTestPanel(false)}
|
||||
/>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
export { createDefaultStep };
|
||||
export default PipelineFlowBuilder;
|
||||
|
|
@ -3,6 +3,7 @@ import { Card, Badge, Button } from "@tremor/react";
|
|||
import { ArrowLeftIcon, PencilIcon } from "@heroicons/react/outline";
|
||||
import { Descriptions, Tag, Spin, Divider, Typography, Alert } from "antd";
|
||||
import { Policy } from "./types";
|
||||
import { PipelineInfoDisplay } from "./pipeline_flow_builder";
|
||||
import { getResolvedGuardrails } from "../networking";
|
||||
|
||||
const { Title, Text } = Typography;
|
||||
|
|
@ -127,6 +128,21 @@ const PolicyInfoView: React.FC<PolicyInfoViewProps> = ({
|
|||
</Descriptions.Item>
|
||||
</Descriptions>
|
||||
|
||||
{policy.pipeline && (
|
||||
<>
|
||||
<Divider orientation="left">
|
||||
<Text strong>Pipeline Flow</Text>
|
||||
</Divider>
|
||||
<Alert
|
||||
message={`Pipeline (${policy.pipeline.mode} mode, ${policy.pipeline.steps.length} step${policy.pipeline.steps.length !== 1 ? "s" : ""})`}
|
||||
type="info"
|
||||
showIcon
|
||||
style={{ marginBottom: 16 }}
|
||||
/>
|
||||
<PipelineInfoDisplay pipeline={policy.pipeline} />
|
||||
</>
|
||||
)}
|
||||
|
||||
<Divider orientation="left">
|
||||
<Text strong>Guardrails Configuration</Text>
|
||||
</Divider>
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ export interface Policy {
|
|||
guardrails_add: string[];
|
||||
guardrails_remove: string[];
|
||||
condition: PolicyCondition | null;
|
||||
pipeline?: GuardrailPipeline | null;
|
||||
created_at?: string;
|
||||
updated_at?: string;
|
||||
created_by?: string;
|
||||
|
|
@ -16,6 +17,19 @@ export interface PolicyCondition {
|
|||
model?: string;
|
||||
}
|
||||
|
||||
export interface PipelineStep {
|
||||
guardrail: string;
|
||||
on_fail: "block" | "allow" | "next" | "modify_response";
|
||||
on_pass: "allow" | "block" | "next" | "modify_response";
|
||||
pass_data?: boolean;
|
||||
modify_response_message?: string | null;
|
||||
}
|
||||
|
||||
export interface GuardrailPipeline {
|
||||
mode: "pre_call" | "post_call";
|
||||
steps: PipelineStep[];
|
||||
}
|
||||
|
||||
export interface PolicyAttachment {
|
||||
attachment_id: string;
|
||||
policy_name: string;
|
||||
|
|
@ -37,6 +51,7 @@ export interface PolicyCreateRequest {
|
|||
guardrails_add?: string[];
|
||||
guardrails_remove?: string[];
|
||||
condition?: PolicyCondition;
|
||||
pipeline?: GuardrailPipeline | null;
|
||||
}
|
||||
|
||||
export interface PolicyUpdateRequest {
|
||||
|
|
@ -46,6 +61,7 @@ export interface PolicyUpdateRequest {
|
|||
guardrails_add?: string[];
|
||||
guardrails_remove?: string[];
|
||||
condition?: PolicyCondition;
|
||||
pipeline?: GuardrailPipeline | null;
|
||||
}
|
||||
|
||||
export interface PolicyAttachmentCreateRequest {
|
||||
|
|
@ -66,3 +82,20 @@ export interface PolicyAttachmentListResponse {
|
|||
attachments: PolicyAttachment[];
|
||||
total_count: number;
|
||||
}
|
||||
|
||||
export interface PipelineStepResult {
|
||||
guardrail_name: string;
|
||||
outcome: "pass" | "fail" | "error";
|
||||
action_taken: string;
|
||||
modified_data: Record<string, any> | null;
|
||||
error_detail: string | null;
|
||||
duration_seconds: number | null;
|
||||
}
|
||||
|
||||
export interface PipelineTestResult {
|
||||
terminal_action: string;
|
||||
step_results: PipelineStepResult[];
|
||||
modified_data: Record<string, any> | null;
|
||||
error_message: string | null;
|
||||
modify_response_message: string | null;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@
|
|||
|
||||
import { ConfigType, GeneralSettingsFieldName, useDeleteProxyConfigField, useProxyConfig } from "@/app/(dashboard)/hooks/proxyConfig/useProxyConfig";
|
||||
import { StoreRequestInSpendLogsParams, useStoreRequestInSpendLogs } from "@/app/(dashboard)/hooks/storeRequestInSpendLogs/useStoreRequestInSpendLogs";
|
||||
import NewBadge from "@/components/common_components/NewBadge";
|
||||
import NotificationsManager from "@/components/molecules/notifications_manager";
|
||||
import { parseErrorMessage } from "@/components/shared/errorUtils";
|
||||
import { ClockCircleOutlined } from "@ant-design/icons";
|
||||
|
|
@ -99,7 +98,7 @@ const SpendLogsSettingsModal: React.FC<SpendLogsSettingsModalProps> = ({ isVisib
|
|||
|
||||
return (
|
||||
<Modal
|
||||
title={<span className="flex gap-2"><Typography.Title level={5}>Spend Logs Settings</Typography.Title><NewBadge /></span>}
|
||||
title={<Typography.Title level={5}>Spend Logs Settings</Typography.Title>}
|
||||
open={isVisible}
|
||||
footer={
|
||||
<Space>
|
||||
|
|
|
|||
|
|
@ -9,7 +9,6 @@ import { Row } from "@tanstack/react-table";
|
|||
import { Switch, Tab, TabGroup, TabList, TabPanel, TabPanels } from "@tremor/react";
|
||||
import { Button, Tooltip } from "antd";
|
||||
import { internalUserRoles } from "../../utils/roles";
|
||||
import NewBadge from "../common_components/NewBadge";
|
||||
import DeletedKeysPage from "../DeletedKeysPage/DeletedKeysPage";
|
||||
import DeletedTeamsPage from "../DeletedTeamsPage/DeletedTeamsPage";
|
||||
import { fetchAllKeyAliases } from "../key_team_helpers/filter_helpers";
|
||||
|
|
@ -514,11 +513,11 @@ export default function SpendLogsTable({
|
|||
<TabPanel>
|
||||
<div className="flex items-center justify-between mb-4">
|
||||
<h1 className="text-xl font-semibold">Request Logs</h1>
|
||||
<NewBadge dot><Button
|
||||
<Button
|
||||
icon={<SettingOutlined />}
|
||||
onClick={() => setIsSpendLogsSettingsModalVisible(true)}
|
||||
title="Spend Logs Settings"
|
||||
/></NewBadge>
|
||||
/>
|
||||
</div>
|
||||
{selectedKeyInfo && selectedKeyIdInfoView && selectedKeyInfo.api_key === selectedKeyIdInfoView ? (
|
||||
<KeyInfoView
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue