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:
jo-nike 2026-02-14 02:08:36 -05:00 • committed by GitHub
parent f1c5e7f30a
commit 495ce34165
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
79 changed files with 6037 additions and 102 deletions

View file

@ -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

View 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.

View file

@ -107,7 +107,8 @@ const sidebars = {
items: [
"proxy/alerting",
"proxy/pagerduty",
"proxy/prometheus"
"proxy/prometheus",
"proxy/pyroscope_profiling"
]
},
{

Binary file not shown.

View file

@ -1,3 +0,0 @@
-- AlterTable
ALTER TABLE "LiteLLM_PolicyAttachmentTable" ADD COLUMN "tags" TEXT[] DEFAULT ARRAY[]::TEXT[];

View file

@ -0,0 +1,3 @@
-- AlterTable
ALTER TABLE "LiteLLM_AccessGroupTable" DROP COLUMN "access_model_ids",
ADD COLUMN "access_model_names" TEXT[] DEFAULT ARRAY[]::TEXT[];

View file

@ -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([])

View file

@ -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==",

View file

@ -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

View file

@ -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(

View file

@ -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,

View file

@ -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,
)

View file

@ -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",
)

View 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]

View file

@ -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(

View file

@ -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)

View file

@ -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,

View 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)

View file

@ -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
# ─────────────────────────────────────────────────────────────────────────────

View file

@ -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

View file

@ -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,

View file

@ -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],

View file

@ -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(

View file

@ -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([])

View file

@ -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

View file

@ -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,

View file

@ -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 (

View file

@ -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

View file

@ -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]

View file

@ -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

View file

@ -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",

View 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

View file

@ -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")

View file

@ -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."""

View file

@ -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
View file

@ -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"]

View file

@ -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"]

View file

@ -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

View file

@ -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([])

View file

@ -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"

View file

@ -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

View file

@ -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
# ---------------------------------------------------------------------------

View file

@ -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 (/)

View 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

View 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

View file

View 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",
)

View file

@ -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

View file

@ -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" });

View file

@ -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();

View file

@ -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);
},
});
};

View file

@ -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();
});
});

View file

@ -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 || ""),
});
};

View file

@ -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 });
},
});
};

View file

@ -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 });
},
});
};

View file

@ -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),
});
},
});
};

View file

@ -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" ? (

View file

@ -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();
});
});

View file

@ -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>
&nbsp;{"by"}&nbsp;
<DefaultProxyAdminTag userId={accessGroup.created_by} />
</Text>
)}
</Descriptions.Item>
<Descriptions.Item label="Last Updated">
{new Date(accessGroup.updated_at).toLocaleString()}
{accessGroup.updated_by && (
<Text>
&nbsp;{"by"}&nbsp;
<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>
);
}

View file

@ -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>
);
}

View file

@ -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>
);
}

View file

@ -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>
);
}

View file

@ -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();
});
});

View file

@ -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>
);
}

View 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
}

View file

@ -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>;
}

View file

@ -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",

View file

@ -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

View file

@ -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",

View file

@ -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"}

View file

@ -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>
);
};

View file

@ -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 &#8594; {ACTION_LABELS[step.on_pass] || step.on_pass}
</span>
<span className="flex items-center gap-1.5">
<FailIcon /> Fail &#8594; {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;

View file

@ -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>

View file

@ -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;
}

View file

@ -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>

View file

@ -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