mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
Merge remote-tracking branch 'origin/main' into litellm_replica_db_opt_in
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> # Conflicts: # tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py
This commit is contained in:
commit
4d64c76a20
71 changed files with 5313 additions and 530 deletions
12
.github/workflows/test-linting.yml
vendored
12
.github/workflows/test-linting.yml
vendored
|
|
@ -180,6 +180,18 @@ jobs:
|
|||
echo "No changed tests/e2e Python files; skipping."
|
||||
fi
|
||||
|
||||
- name: Run the claude_code harness unit tests
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
run: |
|
||||
if ! git diff --name-only --diff-filter=ACMRD "$GATE_BASE_SHA" HEAD -- ':(glob)tests/e2e/claude_code/**/*.py' ':(glob)tests/e2e/*.py' tests/e2e/claude_code/cron_vm/install_claude_code.sh pyproject.toml uv.lock .github/workflows/test-linting.yml | grep -q .; then
|
||||
echo "No changed claude_code harness files; skipping."
|
||||
exit 0
|
||||
fi
|
||||
retry() { "$@" || { sleep 15; "$@"; } || { sleep 30; "$@"; }; }
|
||||
CLAUDE_VERSION="$(retry uv run --no-sync python tests/e2e/claude_code/pr_gate_version_resolver.py)"
|
||||
tests/e2e/claude_code/cron_vm/install_claude_code.sh "$CLAUDE_VERSION" "$RUNNER_TEMP/claude-cli"
|
||||
PATH="$RUNNER_TEMP/claude-cli:$PATH" uv run --no-sync pytest -q --noconftest -o addopts= -o pythonpath=tests/e2e -p no:rerunfailures tests/e2e/claude_code/_*_unit_tests
|
||||
|
||||
- name: Check for circular imports
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
run: |
|
||||
|
|
|
|||
|
|
@ -1697,6 +1697,63 @@
|
|||
"title": "litellm_video_duration_seconds_metric rate",
|
||||
"type": "timeseries"
|
||||
},
|
||||
{
|
||||
"datasource": {
|
||||
"type": "prometheus",
|
||||
"uid": "${DS_PROMETHEUS}"
|
||||
},
|
||||
"description": "Share of the provider's bill LiteLLM captured as spend over the scheduled capture-rate check's window (needs general_settings.spend_capture_rate_check); NaN while no rate is available",
|
||||
"fieldConfig": {
|
||||
"defaults": {
|
||||
"color": {
|
||||
"mode": "palette-classic"
|
||||
},
|
||||
"custom": {
|
||||
"drawStyle": "line",
|
||||
"fillOpacity": 10,
|
||||
"lineWidth": 1,
|
||||
"showPoints": "never",
|
||||
"spanNulls": false
|
||||
},
|
||||
"unit": "percentunit"
|
||||
},
|
||||
"overrides": []
|
||||
},
|
||||
"gridPos": {
|
||||
"h": 8,
|
||||
"w": 12,
|
||||
"x": 12,
|
||||
"y": 107
|
||||
},
|
||||
"id": 111,
|
||||
"options": {
|
||||
"legend": {
|
||||
"calcs": [],
|
||||
"displayMode": "list",
|
||||
"placement": "bottom",
|
||||
"showLegend": true
|
||||
},
|
||||
"tooltip": {
|
||||
"mode": "multi",
|
||||
"sort": "desc"
|
||||
}
|
||||
},
|
||||
"targets": [
|
||||
{
|
||||
"datasource": {
|
||||
"type": "prometheus",
|
||||
"uid": "${DS_PROMETHEUS}"
|
||||
},
|
||||
"editorMode": "code",
|
||||
"expr": "max by (api_provider) (litellm_spend_capture_rate)",
|
||||
"legendFormat": "{{api_provider}}",
|
||||
"range": true,
|
||||
"refId": "A"
|
||||
}
|
||||
],
|
||||
"title": "litellm_spend_capture_rate",
|
||||
"type": "timeseries"
|
||||
},
|
||||
{
|
||||
"collapsed": false,
|
||||
"gridPos": {
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
# LiteLLM All Prometheus Metrics dashboard
|
||||
|
||||
Every `litellm_*` metric family the proxy can expose on `/metrics` (134 families across 95 panels), grouped into rows: proxy traffic, latency, spend and tokens, cache, LLM API deployments, key and team rate limits, budgets, guardrails, MCP, managed files and batches, users and teams, the Redis circuit breaker, the spend log cleanup job, and the `prometheus_system` service callback metrics (per-service latency, request and failure rates, spend update queue sizes). Panel titles are the metric names so you can grep the JSON for the metric you care about
|
||||
Every `litellm_*` metric family the proxy can expose on `/metrics` (136 families across 97 panels), grouped into rows: proxy traffic, latency, spend and tokens, cache, LLM API deployments, key and team rate limits, budgets, guardrails, MCP, managed files and batches, users and teams, the Redis circuit breaker, the spend log cleanup job, and the `prometheus_system` service callback metrics (per-service latency, request and failure rates, spend update queue sizes). Panel titles are the metric names so you can grep the JSON for the metric you care about
|
||||
|
||||
Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/test_litellm/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard
|
||||
|
||||
|
|
|
|||
43
db_scripts/backfill_key_total_spend.sql
Normal file
43
db_scripts/backfill_key_total_spend.sql
Normal file
|
|
@ -0,0 +1,43 @@
|
|||
-- One-shot backfill of LiteLLM_VerificationToken.total_spend (lifetime spend)
|
||||
-- for keys created before the column was introduced in LiteLLM v1.103.0.
|
||||
--
|
||||
-- The column was added with DEFAULT 0 and no backfill, so keys that predate
|
||||
-- the upgrade report lifetime spend below their current period spend. New
|
||||
-- deployments do not need this script: total_spend is updated at request
|
||||
-- time from the moment the release is deployed. Run it only if you want
|
||||
-- pre-upgrade keys to show their historical lifetime spend. It sets lifetime
|
||||
-- spend to at least the current spend on every key, active and archived,
|
||||
-- because current period spend is a valid lower bound on lifetime spend.
|
||||
-- For keys with no budget reset that is already the exact lifetime value;
|
||||
-- for resetting keys it only recovers the current period. It is idempotent:
|
||||
-- it only touches rows where total_spend is below spend, so re-running is a
|
||||
-- no-op. It touches no spend logs and runs in seconds.
|
||||
--
|
||||
-- IMPORTANT caveats before running:
|
||||
--
|
||||
-- 1. Take a backup of the affected tables first:
|
||||
-- pg_dump "$DATABASE_URL" -t '"LiteLLM_VerificationToken"' -t '"LiteLLM_DeletedVerificationToken"' > key_total_spend_backup.sql
|
||||
--
|
||||
-- 2. A key "resets" when its own budget_duration IS NOT NULL, or when its
|
||||
-- budget_id links to a LiteLLM_BudgetTable row whose budget_duration IS
|
||||
-- NOT NULL (a linked budget resets the key's spend each period too). For
|
||||
-- those keys this script only recovers the current period;
|
||||
-- db_scripts/backfill_key_total_spend_from_spend_logs.sql is an optional
|
||||
-- follow-up that rebuilds the earlier periods from LiteLLM_SpendLogs.
|
||||
--
|
||||
-- 3. No proxy restart is needed. The proxy picks up the corrected values on
|
||||
-- its next read of each key.
|
||||
--
|
||||
-- Usage:
|
||||
-- psql "$DATABASE_URL" -f db_scripts/backfill_key_total_spend.sql
|
||||
|
||||
UPDATE "LiteLLM_VerificationToken"
|
||||
SET total_spend = spend
|
||||
WHERE total_spend < spend;
|
||||
|
||||
UPDATE "LiteLLM_DeletedVerificationToken"
|
||||
SET total_spend = spend
|
||||
WHERE total_spend < spend;
|
||||
|
||||
-- Verify: this should return 0.
|
||||
-- SELECT count(*) FROM "LiteLLM_VerificationToken" WHERE total_spend < spend;
|
||||
89
db_scripts/backfill_key_total_spend_from_spend_logs.sql
Normal file
89
db_scripts/backfill_key_total_spend_from_spend_logs.sql
Normal file
|
|
@ -0,0 +1,89 @@
|
|||
-- Optional follow-up to db_scripts/backfill_key_total_spend.sql. Run that
|
||||
-- script first; this one rebuilds earlier budget periods for the keys it
|
||||
-- can only partially fix: keys whose spend resets each period, because their own
|
||||
-- budget_duration IS NOT NULL or because their budget_id links to a
|
||||
-- LiteLLM_BudgetTable row whose budget_duration IS NOT NULL.
|
||||
--
|
||||
-- For those keys the "spend" column only covers the current period, so
|
||||
-- lifetime spend is reconstructed from LiteLLM_SpendLogs. The join matches
|
||||
-- l.api_key against both the stored token and its second sha256
|
||||
-- (encode(sha256(convert_to(token, 'UTF8')), 'hex')), because spend logs
|
||||
-- written by older paths recorded the re-hashed digest instead of the
|
||||
-- token. It is idempotent and never lowers a value: every statement only
|
||||
-- touches rows where total_spend is below the rebuilt sum, so re-running is
|
||||
-- a no-op, and a key whose log history is shorter than its current period
|
||||
-- keeps the value backfill_key_total_spend.sql already gave it.
|
||||
--
|
||||
-- IMPORTANT caveats before running:
|
||||
--
|
||||
-- 1. Take a backup of the affected tables first:
|
||||
-- pg_dump "$DATABASE_URL" -t '"LiteLLM_VerificationToken"' -t '"LiteLLM_DeletedVerificationToken"' > key_total_spend_backup.sql
|
||||
--
|
||||
-- 2. It requires spend logs to have been enabled, and coverage is bounded
|
||||
-- by maximum_spend_logs_retention_period: spend older than the retention
|
||||
-- window is already gone and cannot be recovered.
|
||||
--
|
||||
-- 3. On a large SpendLogs table the join scan is slow, so run it off peak.
|
||||
--
|
||||
-- 4. Run it while the proxy is idle (or with traffic paused). The proxy
|
||||
-- flushes spend logs in batches, so a request that already raised
|
||||
-- total_spend but whose log is still queued is missing from the sum, and
|
||||
-- the rebuilt value would be short by that in-flight amount.
|
||||
--
|
||||
-- 5. A custom token can be deleted and recreated, so the archived table can
|
||||
-- hold several lifetimes of one token. The update only rewrites archived
|
||||
-- rows that reset, and the log sum covers every lifetime of that token.
|
||||
--
|
||||
-- 6. No proxy restart is needed. The proxy picks up the corrected values on
|
||||
-- its next read of each key.
|
||||
--
|
||||
-- Usage:
|
||||
-- psql "$DATABASE_URL" -f db_scripts/backfill_key_total_spend_from_spend_logs.sql
|
||||
|
||||
-- Active keys whose spend resets (own budget_duration, or a linked
|
||||
-- LiteLLM_BudgetTable row with one). Rebuild from LiteLLM_SpendLogs,
|
||||
-- matching api_key against the stored token and its second sha256 digest.
|
||||
UPDATE "LiteLLM_VerificationToken" k
|
||||
SET total_spend = s.sum_spend
|
||||
FROM (
|
||||
SELECT k2.token, SUM(l.spend) AS sum_spend
|
||||
FROM "LiteLLM_VerificationToken" k2
|
||||
JOIN "LiteLLM_SpendLogs" l
|
||||
ON l.api_key IN (k2.token, encode(sha256(convert_to(k2.token, 'UTF8')), 'hex'))
|
||||
WHERE k2.budget_duration IS NOT NULL
|
||||
OR k2.budget_id IN (
|
||||
SELECT budget_id FROM "LiteLLM_BudgetTable" WHERE budget_duration IS NOT NULL
|
||||
)
|
||||
GROUP BY k2.token
|
||||
) s
|
||||
WHERE k.token = s.token
|
||||
AND k.total_spend < s.sum_spend;
|
||||
|
||||
-- Archived tokens are not unique, so collapse them to one row per token
|
||||
-- before joining spend logs; the update then hits every resetting archived
|
||||
-- row.
|
||||
UPDATE "LiteLLM_DeletedVerificationToken" k
|
||||
SET total_spend = s.sum_spend
|
||||
FROM (
|
||||
SELECT k2.token, SUM(l.spend) AS sum_spend
|
||||
FROM (
|
||||
SELECT DISTINCT token
|
||||
FROM "LiteLLM_DeletedVerificationToken"
|
||||
WHERE budget_duration IS NOT NULL
|
||||
OR budget_id IN (
|
||||
SELECT budget_id FROM "LiteLLM_BudgetTable" WHERE budget_duration IS NOT NULL
|
||||
)
|
||||
) k2
|
||||
JOIN "LiteLLM_SpendLogs" l
|
||||
ON l.api_key IN (k2.token, encode(sha256(convert_to(k2.token, 'UTF8')), 'hex'))
|
||||
GROUP BY k2.token
|
||||
) s
|
||||
WHERE k.token = s.token
|
||||
AND k.total_spend < s.sum_spend
|
||||
AND (k.budget_duration IS NOT NULL
|
||||
OR k.budget_id IN (
|
||||
SELECT budget_id FROM "LiteLLM_BudgetTable" WHERE budget_duration IS NOT NULL
|
||||
));
|
||||
|
||||
-- Verify: this should return 0.
|
||||
-- SELECT count(*) FROM "LiteLLM_VerificationToken" WHERE total_spend < spend;
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "kill_switch" JSONB;
|
||||
|
|
@ -72,6 +72,7 @@ model LiteLLM_AgentsTable {
|
|||
agent_card_params Json
|
||||
static_headers Json? @default("{}")
|
||||
extra_headers String[] @default([])
|
||||
kill_switch Json?
|
||||
agent_access_groups String[] @default([])
|
||||
access_group_ids String[] @default([])
|
||||
object_permission_id String?
|
||||
|
|
|
|||
|
|
@ -557,6 +557,8 @@ SEMANTIC_CACHE_EMBEDDING_TIMEOUT_SECONDS: Final[float] = float(
|
|||
request_timeout: float = float(os.getenv("REQUEST_TIMEOUT", str(int(DEFAULT_REQUEST_TIMEOUT_SECONDS))))
|
||||
request_timeout_explicitly_set: bool = "REQUEST_TIMEOUT" in os.environ
|
||||
DEFAULT_A2A_AGENT_TIMEOUT: Final[float] = float(os.getenv("DEFAULT_A2A_AGENT_TIMEOUT", 6000)) # 10 minutes
|
||||
AGENT_KILL_SWITCH_TIMEOUT_SECONDS: Final = 10.0
|
||||
AGENT_KILL_SWITCH_RESPONSE_BODY_MAX_CHARS: Final = 2000
|
||||
# Patterns that indicate a localhost/internal URL in A2A agent cards that should be
|
||||
# replaced with the original base_url. This is a common misconfiguration where
|
||||
# developers deploy agents with development URLs in their agent cards.
|
||||
|
|
@ -2120,6 +2122,14 @@ PTU_LAPSED_ALERT_LIMIT: Final[int] = 10
|
|||
DAILY_GLOBAL_SPEND_RECONCILE_JOB_ID: Final[str] = "daily_global_spend_reconcile_job"
|
||||
DAILY_GLOBAL_SPEND_RECONCILE_LOCK_TTL_SECONDS: Final[int] = 3600
|
||||
DAILY_GLOBAL_SPEND_RECONCILED_THROUGH_PARAM: Final[str] = "daily_global_spend_reconciled_through"
|
||||
SPEND_CAPTURE_RATE_CHECK_JOB_ID: Final[str] = "spend_capture_rate_check_job"
|
||||
SPEND_CAPTURE_RATE_CHECK_LOCK_TTL_SECONDS: Final[int] = 900
|
||||
SPEND_CAPTURE_RATE_MAX_RANGE_DAYS: Final[int] = 180
|
||||
SPEND_CAPTURE_RATE_DOCS_URL: Final[str] = "https://docs.litellm.ai/docs/proxy/spend_capture_rate"
|
||||
OPENAI_ORGANIZATION_COSTS_URL: Final[str] = "https://api.openai.com/v1/organization/costs"
|
||||
# Buckets per page the OpenAI costs endpoint allows (1 to 180, default 7), 2026-09-24
|
||||
OPENAI_ORGANIZATION_COSTS_PAGE_LIMIT: Final[int] = 180
|
||||
PROVIDER_BILLING_TIMEOUT_SECONDS: Final[float] = 30.0
|
||||
# Slack allowed when deciding a sentinel row is stale. The row's updated_at and the
|
||||
# run's cutoff are stamped by different hosts, so clock skew between them must not let
|
||||
# one run delete a charge another just wrote. A stale row is hours old and a concurrent
|
||||
|
|
|
|||
|
|
@ -729,6 +729,15 @@ class PrometheusLogger(CustomLogger):
|
|||
labelnames=self.get_labels_for_metric("litellm_zero_cost_requests_total"),
|
||||
)
|
||||
|
||||
self.litellm_spend_capture_rate = self._gauge_factory(
|
||||
"litellm_spend_capture_rate",
|
||||
(
|
||||
"Share of the provider's bill LiteLLM captured as spend over the scheduled check's window "
|
||||
"(captured spend / provider bill), by api_provider; NaN when the last check produced no rate"
|
||||
),
|
||||
labelnames=self.get_labels_for_metric("litellm_spend_capture_rate"),
|
||||
)
|
||||
|
||||
# Cache metrics
|
||||
self.litellm_cache_hits_metric = self._counter_factory(
|
||||
name="litellm_cache_hits_metric",
|
||||
|
|
@ -2028,6 +2037,15 @@ class PrometheusLogger(CustomLogger):
|
|||
)
|
||||
self.litellm_zero_cost_requests_total.labels(**labels).inc()
|
||||
|
||||
def set_spend_capture_rate(self, api_provider: str, capture_rate: float | None) -> None:
|
||||
labels: Final = prometheus_label_factory(
|
||||
supported_enum_labels=self.get_labels_for_metric("litellm_spend_capture_rate"),
|
||||
enum_values=UserAPIKeyLabelValues(api_provider=api_provider),
|
||||
)
|
||||
gauge: Final = self.litellm_spend_capture_rate
|
||||
series: Final = gauge.labels(**labels) if labels else gauge
|
||||
series.set(math.nan if capture_rate is None else capture_rate)
|
||||
|
||||
@staticmethod
|
||||
def _get_remaining_from_v3_rate_limit_headers(
|
||||
standard_logging_payload: StandardLoggingPayload | None,
|
||||
|
|
|
|||
|
|
@ -385,6 +385,10 @@ def _get_cached_prometheus_logger():
|
|||
return _PrometheusLogger
|
||||
|
||||
|
||||
class RawRequestCaptured(Exception):
|
||||
pass
|
||||
|
||||
|
||||
_DEPLOYMENT_PRICING_KEYS: Final = (
|
||||
"input_cost_per_token",
|
||||
"output_cost_per_token",
|
||||
|
|
@ -591,6 +595,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
kwargs: dict | None = None,
|
||||
log_raw_request_response: bool = False,
|
||||
supports_correlation_logging: bool = True,
|
||||
raw_request_only: bool = False,
|
||||
):
|
||||
_input: Final[str | None] = messages # save original value of messages
|
||||
if messages is not None:
|
||||
|
|
@ -650,6 +655,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
self.streaming_chunks: list[Any] = [] # for generating complete stream response
|
||||
self.sync_streaming_chunks: list[Any] = [] # for generating complete stream response
|
||||
self.log_raw_request_response = log_raw_request_response
|
||||
self.raw_request_only = raw_request_only
|
||||
|
||||
# Initialize dynamic callbacks
|
||||
self.dynamic_input_callbacks: list[str | Callable | CustomLogger] | None = dynamic_input_callbacks
|
||||
|
|
@ -1476,6 +1482,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
if capture_exception: # log this error to sentry for debugging
|
||||
capture_exception(e)
|
||||
|
||||
if self.raw_request_only:
|
||||
raise RawRequestCaptured()
|
||||
|
||||
def _print_llm_call_debugging_log(
|
||||
self,
|
||||
api_base: str,
|
||||
|
|
|
|||
54
litellm/llms/base_llm/files/batch_records.py
Normal file
54
litellm/llms/base_llm/files/batch_records.py
Normal file
|
|
@ -0,0 +1,54 @@
|
|||
from collections.abc import Iterable, Mapping
|
||||
from functools import cache
|
||||
from types import MappingProxyType
|
||||
from typing import Final, cast, get_type_hints
|
||||
|
||||
from litellm.types.llms.openai import ResponseInputParam, ResponsesAPIOptionalRequestParams
|
||||
|
||||
|
||||
def _frozen_mapping(items: Iterable[tuple[str, object]]) -> Mapping[str, object]:
|
||||
return MappingProxyType(dict(items))
|
||||
|
||||
|
||||
@cache
|
||||
def _responses_request_keys() -> frozenset[str]:
|
||||
return frozenset(get_type_hints(ResponsesAPIOptionalRequestParams))
|
||||
|
||||
|
||||
def responses_batch_body_to_chat_body(
|
||||
openai_request_body: Mapping[str, object],
|
||||
custom_llm_provider: str | None = None,
|
||||
) -> dict[str, object]: # mutable-ok: provider transforms take the bridged chat body as a plain dict
|
||||
"""
|
||||
Rewrite the body of an OpenAI `/v1/responses` batch record as a Chat Completions body.
|
||||
|
||||
Batch providers translate chat bodies into their own request shape, so a Responses
|
||||
record goes through the same Responses-to-Chat bridge the real-time path uses for
|
||||
providers without a native Responses API: `input`, `instructions`, `max_output_tokens`
|
||||
and the tool params translate identically in batch and real time. Like real time, the
|
||||
record's fields are forwarded as sent instead of validated against the SDK TypedDicts,
|
||||
whose required keys (a function tool's `strict`, an image part's `detail`) clients omit.
|
||||
"""
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
|
||||
responses_input: Final = openai_request_body.get("input")
|
||||
if responses_input is None:
|
||||
raise ValueError(
|
||||
"Batch record for /v1/responses is missing required `input` field: "
|
||||
f"model={openai_request_body.get('model', '')}"
|
||||
)
|
||||
model: Final = openai_request_body.get("model")
|
||||
chat_input: Final = cast(str | ResponseInputParam, responses_input) # cast-ok: forwarded as sent
|
||||
responses_request: Final = cast( # cast-ok: client-supplied fields forwarded verbatim, as real time does
|
||||
ResponsesAPIOptionalRequestParams,
|
||||
_frozen_mapping((key, value) for key, value in openai_request_body.items() if key in _responses_request_keys()),
|
||||
)
|
||||
return LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request( # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # transformer declares a bare dict return
|
||||
model=model if isinstance(model, str) else "",
|
||||
input=chat_input,
|
||||
responses_api_request=responses_request,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
metadata=openai_request_body.get("metadata"),
|
||||
)
|
||||
|
|
@ -8,7 +8,6 @@ from collections.abc import Iterable, Mapping, MutableMapping, Sequence
|
|||
from contextlib import suppress
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from functools import cache
|
||||
from itertools import chain
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final, Literal, TypeAlias, TypedDict
|
||||
|
|
@ -17,7 +16,7 @@ from urllib.parse import quote, unquote, urlencode
|
|||
import httpx
|
||||
from httpx import Headers, Response
|
||||
from openai.types.file_deleted import FileDeleted
|
||||
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -41,7 +40,9 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
extract_file_data,
|
||||
text_completion_prompt_to_messages,
|
||||
)
|
||||
from litellm.llms.base_llm.base_utils import map_developer_role_to_system_role
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.base_llm.files.batch_records import responses_batch_body_to_chat_body
|
||||
from litellm.llms.base_llm.files.transformation import (
|
||||
BaseFilesConfig,
|
||||
LiteLLMLoggingObj,
|
||||
|
|
@ -56,8 +57,6 @@ from litellm.types.llms.openai import (
|
|||
OpenAICreateFileRequestOptionalParams,
|
||||
OpenAIFileObject,
|
||||
PathLike,
|
||||
ResponseInputParam,
|
||||
ResponsesAPIOptionalRequestParams,
|
||||
)
|
||||
from litellm.types.utils import ExtractedFileData, LlmProviders, SpecialEnums
|
||||
from litellm.utils import get_llm_provider
|
||||
|
|
@ -130,22 +129,6 @@ class _S3UploadResponse(TypedDict, total=False):
|
|||
ContentLength: ReadOnly[int]
|
||||
|
||||
|
||||
# JSONL batch records are untyped json, so the `/v1/responses` fields are
|
||||
# validated into their concrete Responses API types before being handed to the
|
||||
# Responses-to-Chat bridge. Both adapters drop keys the Responses API doesn't
|
||||
# define, which is what the bridge would ignore anyway. Built on first use
|
||||
# rather than at import: `ResponseInputParam` is a deep union and only batch
|
||||
# files carrying `/v1/responses` records need it.
|
||||
@cache
|
||||
def _responses_input_adapter() -> TypeAdapter[str | ResponseInputParam]:
|
||||
return TypeAdapter(str | ResponseInputParam)
|
||||
|
||||
|
||||
@cache
|
||||
def _responses_request_adapter() -> TypeAdapter[ResponsesAPIOptionalRequestParams]:
|
||||
return TypeAdapter(ResponsesAPIOptionalRequestParams)
|
||||
|
||||
|
||||
class _BedrockS3RequestParams(AwsAuthParams):
|
||||
"""Typed view of the credential/region params the S3 GetObject path reads."""
|
||||
|
||||
|
|
@ -859,33 +842,9 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
Delegates to the same Responses-to-Chat bridge the real-time path uses
|
||||
for providers without a native Responses API (which is every Bedrock
|
||||
model), so `input`, `instructions`, `max_output_tokens` and the tool
|
||||
params translate identically in batch and real time. The bridge always
|
||||
emits a `tools` key; an empty one is dropped rather than shipped as an
|
||||
empty array inside `modelInput`.
|
||||
params translate identically in batch and real time.
|
||||
"""
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
|
||||
responses_input: Final = openai_request_body.get("input")
|
||||
if responses_input is None:
|
||||
raise ValueError(
|
||||
"Batch record for /v1/responses is missing required `input` field: "
|
||||
f"model={openai_request_body.get('model', '')}"
|
||||
)
|
||||
chat_body: Final[Mapping[str, object]] = (
|
||||
LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request(
|
||||
model=openai_request_body.get("model", ""),
|
||||
input=_responses_input_adapter().validate_python(responses_input),
|
||||
responses_api_request=_responses_request_adapter().validate_python(
|
||||
_frozen_mapping(
|
||||
(key, value) for key, value in openai_request_body.items() if key not in ("model", "input")
|
||||
)
|
||||
),
|
||||
metadata=openai_request_body.get("metadata"),
|
||||
)
|
||||
)
|
||||
return _frozen_mapping((key, value) for key, value in chat_body.items() if key != "tools" or value)
|
||||
return responses_batch_body_to_chat_body(openai_request_body)
|
||||
|
||||
@staticmethod
|
||||
def _transform_batch_body_to_chat_body(
|
||||
|
|
@ -922,7 +881,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
"""
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
messages: Final = openai_request_body.get("messages", [])
|
||||
messages: Final = map_developer_role_to_system_role(openai_request_body.get("messages", []))
|
||||
optional_params: Final = {k: v for k, v in openai_request_body.items() if k not in ["model", "messages"]}
|
||||
|
||||
# --- Anthropic: use existing AmazonAnthropicClaudeConfig ---
|
||||
|
|
|
|||
133
litellm/llms/openai/organization_costs.py
Normal file
133
litellm/llms/openai/organization_costs.py
Normal file
|
|
@ -0,0 +1,133 @@
|
|||
"""OpenAI's organization costs endpoint: the USD the organization was billed per UTC day, read with an admin key."""
|
||||
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import date, datetime, timedelta, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, ValidationError
|
||||
|
||||
from litellm.constants import (
|
||||
OPENAI_ORGANIZATION_COSTS_PAGE_LIMIT,
|
||||
OPENAI_ORGANIZATION_COSTS_URL,
|
||||
PROVIDER_BILLING_TIMEOUT_SECONDS,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
||||
OPENAI_ADMIN_KEY_ENV_VAR: Final = "OPENAI_ADMIN_KEY"
|
||||
|
||||
BillingHttpGet: TypeAlias = Callable[
|
||||
[str, Mapping[str, object], Mapping[str, str]], # mutable-ok: Callable parameter list is type syntax
|
||||
Awaitable[httpx.Response],
|
||||
]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class OpenAICostsRequestFailed:
|
||||
detail: str
|
||||
|
||||
|
||||
class _OpenAICostAmount(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
value: float
|
||||
currency: Literal["usd"]
|
||||
|
||||
|
||||
class _OpenAICostResult(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
amount: _OpenAICostAmount
|
||||
|
||||
|
||||
class _OpenAICostBucket(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
start_time: int
|
||||
results: tuple[_OpenAICostResult, ...] = ()
|
||||
|
||||
|
||||
class _OpenAICostsPage(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
data: tuple[_OpenAICostBucket, ...]
|
||||
has_more: bool = False
|
||||
next_page: str | None = None
|
||||
|
||||
|
||||
async def provider_billing_get(url: str, params: Mapping[str, object], headers: Mapping[str, str]) -> httpx.Response:
|
||||
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.ProviderBilling)
|
||||
return await client.get(
|
||||
url,
|
||||
params=dict(params), # mutable-ok: AsyncHTTPHandler.get takes dict params
|
||||
headers=dict(headers), # mutable-ok: AsyncHTTPHandler.get takes dict headers
|
||||
timeout=PROVIDER_BILLING_TIMEOUT_SECONDS,
|
||||
)
|
||||
|
||||
|
||||
def _utc_midnight(day: date) -> int:
|
||||
return int(datetime(day.year, day.month, day.day, tzinfo=timezone.utc).timestamp())
|
||||
|
||||
|
||||
def _bucket_day(bucket: _OpenAICostBucket) -> str:
|
||||
return datetime.fromtimestamp(bucket.start_time, tz=timezone.utc).date().isoformat()
|
||||
|
||||
|
||||
async def fetch_openai_daily_costs(
|
||||
start_date: date,
|
||||
end_date: date,
|
||||
*,
|
||||
admin_key: str,
|
||||
project_ids: Sequence[str] = (),
|
||||
http_get: BillingHttpGet = provider_billing_get,
|
||||
) -> Mapping[str, float] | OpenAICostsRequestFailed:
|
||||
"""USD billed by OpenAI per UTC day (ISO date) over the closed range, following pagination to the end."""
|
||||
scope: Final = (("project_ids[]", tuple(project_ids)),) if project_ids else ()
|
||||
window: Final[Mapping[str, object]] = MappingProxyType(
|
||||
{
|
||||
key: value
|
||||
for key, value in (
|
||||
("start_time", _utc_midnight(start_date)),
|
||||
("end_time", _utc_midnight(end_date + timedelta(days=1))),
|
||||
("bucket_width", "1d"),
|
||||
("limit", OPENAI_ORGANIZATION_COSTS_PAGE_LIMIT),
|
||||
*scope,
|
||||
)
|
||||
}
|
||||
)
|
||||
headers: Final[Mapping[str, str]] = MappingProxyType({"Authorization": f"Bearer {admin_key}"})
|
||||
|
||||
async def fetch_from(page: str | None) -> tuple[_OpenAICostBucket, ...] | OpenAICostsRequestFailed:
|
||||
params: Final[Mapping[str, object]] = MappingProxyType(
|
||||
{key: value for key, value in (*window.items(), ("page", page)) if value is not None}
|
||||
)
|
||||
try:
|
||||
response: Final = await http_get(OPENAI_ORGANIZATION_COSTS_URL, params, headers)
|
||||
except httpx.HTTPError as exc:
|
||||
return OpenAICostsRequestFailed(f"request failed: {exc}")
|
||||
if response.status_code != 200:
|
||||
return OpenAICostsRequestFailed(f"HTTP {response.status_code}: {response.text[:300]}")
|
||||
try:
|
||||
parsed: Final = _OpenAICostsPage.model_validate(response.json())
|
||||
except (ValueError, ValidationError) as exc:
|
||||
return OpenAICostsRequestFailed(f"unexpected response shape: {exc}")
|
||||
if not parsed.has_more or parsed.next_page is None:
|
||||
return parsed.data
|
||||
rest: Final = await fetch_from(parsed.next_page)
|
||||
return rest if isinstance(rest, OpenAICostsRequestFailed) else parsed.data + rest
|
||||
|
||||
buckets: Final = await fetch_from(None)
|
||||
if isinstance(buckets, OpenAICostsRequestFailed):
|
||||
return buckets
|
||||
days: Final = frozenset(_bucket_day(bucket) for bucket in buckets)
|
||||
return MappingProxyType(
|
||||
{
|
||||
day: sum(
|
||||
result.amount.value for bucket in buckets if _bucket_day(bucket) == day for result in bucket.results
|
||||
)
|
||||
for day in days
|
||||
}
|
||||
)
|
||||
|
|
@ -35,7 +35,9 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
extract_file_data,
|
||||
extract_file_metadata,
|
||||
)
|
||||
from litellm.llms.base_llm.base_utils import map_developer_role_to_system_role
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.base_llm.files.batch_records import responses_batch_body_to_chat_body
|
||||
from litellm.llms.base_llm.files.transformation import (
|
||||
BaseFilesConfig,
|
||||
BaseFileUploadStream,
|
||||
|
|
@ -529,21 +531,30 @@ def is_passthrough_batch_upload(create_file_data: Mapping[str, object], litellm_
|
|||
return create_file_data.get("purpose") == "batch" and litellm_params.get("passthrough") is True
|
||||
|
||||
|
||||
def _is_embeddings_batch_entry(openai_entry: Mapping[str, object]) -> bool:
|
||||
def _batch_entry_route_path(openai_entry: Mapping[str, object]) -> str:
|
||||
"""
|
||||
Whether an OpenAI batch JSONL line targets the embeddings endpoint.
|
||||
The route an OpenAI batch JSONL line targets, without query string or trailing slash.
|
||||
|
||||
OpenAI puts the target route on each line's `url` (e.g. `/v1/embeddings`); Vertex
|
||||
has no equivalent per-line field, so the route decides which Vertex request shape
|
||||
the line has to be translated into.
|
||||
"""
|
||||
url = openai_entry.get("url")
|
||||
url: Final = openai_entry.get("url")
|
||||
if not isinstance(url, str):
|
||||
return False
|
||||
path = url.split("?")[0].rstrip("/")
|
||||
return ""
|
||||
return url.split("?")[0].rstrip("/")
|
||||
|
||||
|
||||
def _is_embeddings_batch_entry(openai_entry: Mapping[str, object]) -> bool:
|
||||
path: Final = _batch_entry_route_path(openai_entry)
|
||||
return path == "embeddings" or path.endswith("/embeddings")
|
||||
|
||||
|
||||
def _is_responses_batch_entry(openai_entry: Mapping[str, object]) -> bool:
|
||||
path: Final = _batch_entry_route_path(openai_entry)
|
||||
return path == "responses" or path.endswith("/responses")
|
||||
|
||||
|
||||
def _openai_embedding_input_elements(
|
||||
embedding_input: GeminiEmbeddingInput,
|
||||
) -> tuple[str | list[str], ...]:
|
||||
|
|
@ -665,10 +676,15 @@ def _openai_batch_jsonl_entry_to_vertex_rows(
|
|||
return _openai_batch_jsonl_entry_to_vertex_embeddings_rows(openai_entry)
|
||||
|
||||
openai_request_body: Final = openai_entry.get("body") or {}
|
||||
chat_request_body: Final = (
|
||||
responses_batch_body_to_chat_body(openai_request_body, custom_llm_provider="vertex_ai")
|
||||
if _is_responses_batch_entry(openai_entry)
|
||||
else openai_request_body
|
||||
)
|
||||
vertex_request_body: Final = _transform_request_body(
|
||||
messages=openai_request_body.get("messages", []),
|
||||
model=openai_request_body.get("model", ""),
|
||||
optional_params=map_openai_to_vertex_params(openai_request_body),
|
||||
messages=map_developer_role_to_system_role(chat_request_body.get("messages", [])),
|
||||
model=chat_request_body.get("model", ""),
|
||||
optional_params=map_openai_to_vertex_params(chat_request_body),
|
||||
custom_llm_provider="vertex_ai",
|
||||
litellm_params={},
|
||||
cached_content=None,
|
||||
|
|
|
|||
|
|
@ -26207,7 +26207,6 @@
|
|||
},
|
||||
"gemini-2.5-flash-image": {
|
||||
"deprecation_date": "2027-03-15",
|
||||
"cache_read_input_token_cost": 3e-08,
|
||||
"input_cost_per_audio_token": 1e-06,
|
||||
"input_cost_per_token": 3e-07,
|
||||
"input_cost_per_token_batches": 1.5e-07,
|
||||
|
|
@ -26312,10 +26311,14 @@
|
|||
"gemini-3-pro-image-preview": {
|
||||
"input_cost_per_image": 0.0011,
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_priority": 3.6e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07,
|
||||
"cache_read_input_token_cost_batches": 1e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_priority": 3.6e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_above_200k_tokens_priority": 7.2e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 65536,
|
||||
|
|
@ -26325,7 +26328,9 @@
|
|||
"output_cost_per_image": 0.134,
|
||||
"output_cost_per_image_token": 0.00012,
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_priority": 2.16e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_above_200k_tokens_priority": 3.24e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/pricing",
|
||||
"supported_endpoints": [
|
||||
|
|
@ -27848,6 +27853,7 @@
|
|||
"gemini-embedding-001": {
|
||||
"deprecation_date": "2028-05-20",
|
||||
"input_cost_per_token": 1.5e-07,
|
||||
"input_cost_per_token_batches": 1.2e-07,
|
||||
"litellm_provider": "vertex_ai-embedding-models",
|
||||
"max_input_tokens": 2048,
|
||||
"max_tokens": 2048,
|
||||
|
|
@ -49410,7 +49416,6 @@
|
|||
},
|
||||
"vertex_ai/gemini-2.5-flash-image": {
|
||||
"deprecation_date": "2027-03-15",
|
||||
"cache_read_input_token_cost": 3e-08,
|
||||
"input_cost_per_audio_token": 1e-06,
|
||||
"input_cost_per_token": 3e-07,
|
||||
"input_cost_per_token_batches": 1.5e-07,
|
||||
|
|
@ -49492,10 +49497,14 @@
|
|||
"vertex_ai/gemini-3-pro-image-preview": {
|
||||
"input_cost_per_image": 0.0011,
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_priority": 3.6e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07,
|
||||
"cache_read_input_token_cost_batches": 1e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_priority": 3.6e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_above_200k_tokens_priority": 7.2e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 65536,
|
||||
|
|
@ -49505,7 +49514,9 @@
|
|||
"output_cost_per_image": 0.134,
|
||||
"output_cost_per_image_token": 0.00012,
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_priority": 2.16e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_above_200k_tokens_priority": 3.24e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"supports_reasoning": false,
|
||||
"source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image"
|
||||
|
|
@ -69473,6 +69484,7 @@
|
|||
"input_cost_per_audio_token": 3e-06,
|
||||
"input_cost_per_image_token": 1e-06,
|
||||
"input_cost_per_token": 7.5e-07,
|
||||
"input_cost_per_video_token": 1e-06,
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 131072,
|
||||
"max_output_tokens": 65536,
|
||||
|
|
@ -69493,6 +69505,7 @@
|
|||
"input_cost_per_audio_token": 3e-06,
|
||||
"input_cost_per_image_token": 1e-06,
|
||||
"input_cost_per_token": 7.5e-07,
|
||||
"input_cost_per_video_token": 1e-06,
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 131072,
|
||||
"max_output_tokens": 65536,
|
||||
|
|
|
|||
|
|
@ -2057,7 +2057,7 @@ class MCPRequestHandler:
|
|||
@staticmethod
|
||||
async def _get_team_object_permission(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
):
|
||||
) -> LiteLLM_ObjectPermissionTable | None:
|
||||
"""
|
||||
Get team object_permission - automatically loaded by get_team_object() in main auth flow.
|
||||
|
||||
|
|
@ -2289,6 +2289,10 @@ class MCPRequestHandler:
|
|||
)
|
||||
)
|
||||
|
||||
allowed_tools = _as_list(
|
||||
await MCPRequestHandler._apply_agent_caller_tool_ceiling(allowed_tools, server_id, user_api_key_auth)
|
||||
)
|
||||
|
||||
return await MCPRequestHandler._apply_agent_and_org_tool_ceilings(
|
||||
allowed_tools, server_id, user_api_key_auth, keyless_source=keyless_source
|
||||
)
|
||||
|
|
@ -3170,6 +3174,48 @@ class MCPRequestHandler:
|
|||
return list(user_tools)
|
||||
return list(set(allowed_tools) & set(user_tools))
|
||||
|
||||
@staticmethod
|
||||
async def _apply_agent_caller_tool_ceiling(
|
||||
allowed_tools: Sequence[str] | None,
|
||||
server_id: str,
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
) -> Sequence[str] | None:
|
||||
"""Narrow an agent key's tools on ``server_id`` to those the invoking user and team (echoed back
|
||||
by the agent as ``x-litellm-user-id`` / ``x-litellm-team-id``) may call: the echoed team's tool
|
||||
grants when it names any on this server, then the echoed user's own tool entitlement. The tools
|
||||
axis twin of ``_apply_agent_caller_ceiling``, so the headers only ever narrow. Denies every tool
|
||||
on the server when the caller's team cannot be loaded, since a caller we cannot resolve must not
|
||||
read as unrestricted."""
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
caller_auth: Final = agent_caller_auth(user_api_key_auth) if user_api_key_auth else None
|
||||
if caller_auth is None:
|
||||
return allowed_tools
|
||||
try:
|
||||
team_obj_perm: Final = await MCPRequestHandler._get_team_object_permission(caller_auth)
|
||||
team_toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(team_obj_perm, server_id)
|
||||
except Exception as e: # noqa: BLE001 # an unresolved caller team must deny, not widen
|
||||
verbose_logger.warning(
|
||||
"MCP agent caller team tool ceiling unresolvable, denying tools on %r: %s", server_id, e
|
||||
)
|
||||
return ()
|
||||
team_direct_tools: Final = (
|
||||
global_mcp_server_manager.expand_tool_permissions(team_obj_perm.mcp_tool_permissions).get(server_id)
|
||||
if team_obj_perm
|
||||
else None
|
||||
)
|
||||
team_tools: Final = MCPRequestHandler._union_tool_grants(team_direct_tools, team_toolset_tools)
|
||||
team_capped: Final = (
|
||||
allowed_tools
|
||||
if team_tools is None
|
||||
else tuple(team_tools)
|
||||
if allowed_tools is None
|
||||
else tuple(frozenset(allowed_tools) & frozenset(team_tools))
|
||||
)
|
||||
return await MCPRequestHandler._apply_user_tool_ceiling(team_capped, server_id, caller_auth)
|
||||
|
||||
@staticmethod
|
||||
async def _apply_end_user_tool_ceiling(
|
||||
allowed_tools: Sequence[str] | None,
|
||||
|
|
|
|||
|
|
@ -2392,6 +2392,16 @@
|
|||
],
|
||||
"title": "Extra Headers"
|
||||
},
|
||||
"kill_switch": {
|
||||
"anyOf": [
|
||||
{
|
||||
"$ref": "#/components/schemas/AgentKillSwitchConfig"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
]
|
||||
},
|
||||
"litellm_params": {
|
||||
"additionalProperties": true,
|
||||
"title": "Litellm Params",
|
||||
|
|
@ -2561,6 +2571,221 @@
|
|||
"title": "AgentKeySummary",
|
||||
"type": "object"
|
||||
},
|
||||
"AgentKillSwitchApiKeyAuth": {
|
||||
"additionalProperties": false,
|
||||
"properties": {
|
||||
"api_key": {
|
||||
"title": "Api Key",
|
||||
"type": "string"
|
||||
},
|
||||
"header_name": {
|
||||
"default": "x-api-key",
|
||||
"title": "Header Name",
|
||||
"type": "string"
|
||||
},
|
||||
"type": {
|
||||
"const": "api_key",
|
||||
"title": "Type",
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"type",
|
||||
"api_key"
|
||||
],
|
||||
"title": "AgentKillSwitchApiKeyAuth",
|
||||
"type": "object"
|
||||
},
|
||||
"AgentKillSwitchBasicAuth": {
|
||||
"additionalProperties": false,
|
||||
"properties": {
|
||||
"password": {
|
||||
"title": "Password",
|
||||
"type": "string"
|
||||
},
|
||||
"type": {
|
||||
"const": "basic",
|
||||
"title": "Type",
|
||||
"type": "string"
|
||||
},
|
||||
"username": {
|
||||
"title": "Username",
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"type",
|
||||
"username",
|
||||
"password"
|
||||
],
|
||||
"title": "AgentKillSwitchBasicAuth",
|
||||
"type": "object"
|
||||
},
|
||||
"AgentKillSwitchBearerAuth": {
|
||||
"additionalProperties": false,
|
||||
"properties": {
|
||||
"token": {
|
||||
"title": "Token",
|
||||
"type": "string"
|
||||
},
|
||||
"type": {
|
||||
"const": "bearer",
|
||||
"title": "Type",
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"type",
|
||||
"token"
|
||||
],
|
||||
"title": "AgentKillSwitchBearerAuth",
|
||||
"type": "object"
|
||||
},
|
||||
"AgentKillSwitchConfig": {
|
||||
"additionalProperties": false,
|
||||
"description": "Webhook an admin fires to shut an agent down out of band. LiteLLM only\nmakes the call; whatever the endpoint does with it is the agent's business.",
|
||||
"properties": {
|
||||
"auth": {
|
||||
"anyOf": [
|
||||
{
|
||||
"discriminator": {
|
||||
"mapping": {
|
||||
"api_key": "#/components/schemas/AgentKillSwitchApiKeyAuth",
|
||||
"basic": "#/components/schemas/AgentKillSwitchBasicAuth",
|
||||
"bearer": "#/components/schemas/AgentKillSwitchBearerAuth"
|
||||
},
|
||||
"propertyName": "type"
|
||||
},
|
||||
"oneOf": [
|
||||
{
|
||||
"$ref": "#/components/schemas/AgentKillSwitchBearerAuth"
|
||||
},
|
||||
{
|
||||
"$ref": "#/components/schemas/AgentKillSwitchApiKeyAuth"
|
||||
},
|
||||
{
|
||||
"$ref": "#/components/schemas/AgentKillSwitchBasicAuth"
|
||||
}
|
||||
]
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Auth"
|
||||
},
|
||||
"body": {
|
||||
"anyOf": [
|
||||
{
|
||||
"additionalProperties": true,
|
||||
"type": "object"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Body"
|
||||
},
|
||||
"headers": {
|
||||
"additionalProperties": {
|
||||
"type": "string"
|
||||
},
|
||||
"title": "Headers",
|
||||
"type": "object"
|
||||
},
|
||||
"method": {
|
||||
"default": "POST",
|
||||
"enum": [
|
||||
"POST",
|
||||
"PUT",
|
||||
"PATCH",
|
||||
"DELETE",
|
||||
"GET"
|
||||
],
|
||||
"title": "Method",
|
||||
"type": "string"
|
||||
},
|
||||
"query_params": {
|
||||
"additionalProperties": {
|
||||
"type": "string"
|
||||
},
|
||||
"title": "Query Params",
|
||||
"type": "object"
|
||||
},
|
||||
"url": {
|
||||
"title": "Url",
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"url"
|
||||
],
|
||||
"title": "AgentKillSwitchConfig",
|
||||
"type": "object"
|
||||
},
|
||||
"AgentKillSwitchResult": {
|
||||
"properties": {
|
||||
"agent_id": {
|
||||
"title": "Agent Id",
|
||||
"type": "string"
|
||||
},
|
||||
"error": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Error"
|
||||
},
|
||||
"method": {
|
||||
"enum": [
|
||||
"POST",
|
||||
"PUT",
|
||||
"PATCH",
|
||||
"DELETE",
|
||||
"GET"
|
||||
],
|
||||
"title": "Method",
|
||||
"type": "string"
|
||||
},
|
||||
"response_body": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Response Body"
|
||||
},
|
||||
"status_code": {
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "integer"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Status Code"
|
||||
},
|
||||
"url": {
|
||||
"title": "Url",
|
||||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"agent_id",
|
||||
"url",
|
||||
"method"
|
||||
],
|
||||
"title": "AgentKillSwitchResult",
|
||||
"type": "object"
|
||||
},
|
||||
"AgentMakePublicResponse": {
|
||||
"properties": {
|
||||
"message": {
|
||||
|
|
@ -2775,6 +3000,16 @@
|
|||
],
|
||||
"title": "Keys"
|
||||
},
|
||||
"kill_switch": {
|
||||
"anyOf": [
|
||||
{
|
||||
"$ref": "#/components/schemas/AgentKillSwitchConfig"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
]
|
||||
},
|
||||
"litellm_params": {
|
||||
"anyOf": [
|
||||
{
|
||||
|
|
@ -3569,6 +3804,16 @@
|
|||
],
|
||||
"title": "Extra Headers"
|
||||
},
|
||||
"kill_switch": {
|
||||
"anyOf": [
|
||||
{
|
||||
"$ref": "#/components/schemas/AgentKillSwitchConfig"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
]
|
||||
},
|
||||
"litellm_params": {
|
||||
"additionalProperties": true,
|
||||
"title": "Litellm Params",
|
||||
|
|
@ -4331,6 +4576,54 @@
|
|||
]
|
||||
}
|
||||
},
|
||||
"/v1/agents/{agent_id}/kill_switch": {
|
||||
"post": {
|
||||
"description": "Fire the agent's configured kill switch webhook. Proxy admin only.\n\nLiteLLM only makes the configured HTTP call and reports what came back; it\ndoes not change the agent's state in LiteLLM. Returns 200 when the webhook\nanswered 2xx, 502 with the same result body otherwise. Every attempt is\nwritten to the audit log as a `kill_switch_fired` row against the agent.\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000/kill_switch\" \\\n -H \"Authorization: Bearer <your_api_key>\"\n```",
|
||||
"operationId": "trigger_agent_kill_switch_v1_agents__agent_id__kill_switch_post",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
"name": "agent_id",
|
||||
"required": true,
|
||||
"schema": {
|
||||
"title": "Agent Id",
|
||||
"type": "string"
|
||||
}
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/AgentKillSwitchResult"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
},
|
||||
"422": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/HTTPValidationError"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Validation Error"
|
||||
}
|
||||
},
|
||||
"security": [
|
||||
{
|
||||
"APIKeyHeader": []
|
||||
}
|
||||
],
|
||||
"summary": "Trigger Agent Kill Switch",
|
||||
"tags": [
|
||||
"agents"
|
||||
]
|
||||
}
|
||||
},
|
||||
"/v1/agents/{agent_id}/make_public": {
|
||||
"post": {
|
||||
"description": "Make an agent publicly discoverable\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000/make_public\" \\\n -H \"Authorization: Bearer <your_api_key>\" \\\n -H \"Content-Type: application/json\"\n```\n\nExample Response:\n```json\n{\n \"agent_id\": \"123e4567-e89b-12d3-a456-426614174000\",\n \"agent_name\": \"my-custom-agent\",\n \"litellm_params\": {\n \"make_public\": true\n },\n \"agent_card_params\": {...},\n \"created_at\": \"2025-11-15T10:30:00Z\",\n \"updated_at\": \"2025-11-15T10:35:00Z\",\n \"created_by\": \"user123\",\n \"updated_by\": \"user123\"\n}\n```",
|
||||
|
|
|
|||
|
|
@ -51,6 +51,7 @@ from litellm.types.proxy.carried_budget_state import (
|
|||
UserBudgetSnapshot,
|
||||
)
|
||||
from litellm.types.proxy.control_plane_endpoints import WorkerRegistryEntry
|
||||
from litellm.types.proxy.spend_capture_rate import SpendCaptureRateCheckSettings
|
||||
from litellm.types.router import RouterErrors, UpdateRouterConfig
|
||||
from litellm.types.router_weights import validate_router_settings_dict
|
||||
from litellm.types.secret_managers.main import KeyManagementSystem
|
||||
|
|
@ -238,6 +239,7 @@ class LitellmTableNames(str, enum.Enum):
|
|||
CONFIG_TABLE_NAME = "LiteLLM_Config"
|
||||
SSO_CONFIG_TABLE_NAME = "LiteLLM_SSOConfig"
|
||||
UI_SETTINGS_TABLE_NAME = "LiteLLM_UISettings"
|
||||
AGENT_TABLE_NAME = "LiteLLM_AgentsTable"
|
||||
|
||||
|
||||
class Litellm_EntityType(enum.Enum):
|
||||
|
|
@ -578,6 +580,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/v1/agents/{agent_id}",
|
||||
"/v1/agents/make_public",
|
||||
"/v1/agents/{agent_id}/make_public",
|
||||
"/v1/agents/{agent_id}/kill_switch",
|
||||
)
|
||||
|
||||
# Backwards-compat union — virtual keys may be configured with
|
||||
|
|
@ -777,6 +780,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/global/spend/provider",
|
||||
"/global/spend/tags",
|
||||
"/global/spend/all_tag_names",
|
||||
"/spend/capture_rate",
|
||||
]
|
||||
|
||||
public_routes = frozenset(
|
||||
|
|
@ -2946,6 +2950,14 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
"every replica. On by default; set to tune the window, pin a job, or turn it off."
|
||||
),
|
||||
)
|
||||
spend_capture_rate_check: SpendCaptureRateCheckSettings | None = Field(
|
||||
None,
|
||||
description=(
|
||||
"Daily check of the spend LiteLLM captured against the provider's own bill (OpenAI via OPENAI_ADMIN_KEY). "
|
||||
"Publishes litellm_spend_capture_rate per provider and alerts when the ratio over the lookback window "
|
||||
"falls under the threshold (default 0.9). Off unless set."
|
||||
),
|
||||
)
|
||||
maximum_spend_logs_retention_period: str | None = Field(
|
||||
None,
|
||||
description="Maximum retention period for spend logs (e.g., '7d' for 7 days). Logs older than this will be deleted.",
|
||||
|
|
@ -3688,7 +3700,7 @@ from litellm.models.spend_logs import ( # noqa: E402
|
|||
)
|
||||
from litellm.models.tag import LiteLLM_TagTable as LiteLLM_TagTable # noqa: E402
|
||||
|
||||
AUDIT_ACTIONS = Literal["created", "updated", "deleted", "blocked", "unblocked", "rotated"]
|
||||
AUDIT_ACTIONS = Literal["created", "updated", "deleted", "blocked", "unblocked", "rotated", "kill_switch_fired"]
|
||||
|
||||
|
||||
class LiteLLM_AuditLogs(LiteLLMPydanticObjectBase):
|
||||
|
|
|
|||
|
|
@ -13,13 +13,14 @@ import litellm
|
|||
from litellm.constants import REDACTED_BY_LITELM_STRING
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
|
||||
from litellm.proxy.agent_endpoints.kill_switch import restore_kill_switch
|
||||
from litellm.proxy.management_helpers.object_permission_utils import (
|
||||
handle_update_object_permission_common,
|
||||
)
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
from litellm.repositories.prisma_protocols import TableActions
|
||||
from litellm.repositories.table_repositories import AgentsRepository, ObjectPermissionRepository
|
||||
from litellm.types.agents import AgentConfig, AgentResponse, PatchAgentRequest
|
||||
from litellm.types.agents import AgentConfig, AgentKillSwitchConfig, AgentResponse, PatchAgentRequest
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import models as prisma_models
|
||||
|
|
@ -31,6 +32,10 @@ class AgentObjectPermissionRecord(Protocol):
|
|||
def dict(self) -> dict[str, object]: ...
|
||||
|
||||
|
||||
class AgentIdWhere(TypedDict):
|
||||
agent_id: ReadOnly[str]
|
||||
|
||||
|
||||
class AgentRecordDump(TypedDict):
|
||||
agent_id: str
|
||||
agent_name: str
|
||||
|
|
@ -38,6 +43,7 @@ class AgentRecordDump(TypedDict):
|
|||
agent_card_params: dict[str, object]
|
||||
static_headers: dict[str, str] | None
|
||||
extra_headers: list[str] | None
|
||||
kill_switch: ReadOnly[AgentKillSwitchConfig | None]
|
||||
access_group_ids: ReadOnly[Sequence[str] | None]
|
||||
object_permission: dict[str, object] | None
|
||||
spend: float
|
||||
|
|
@ -70,6 +76,9 @@ class AgentRecord(Protocol):
|
|||
@property
|
||||
def access_group_ids(self) -> Sequence[str] | None: ...
|
||||
|
||||
@property
|
||||
def kill_switch(self) -> Mapping[str, object] | None: ...
|
||||
|
||||
@property
|
||||
def spend(self) -> float: ...
|
||||
|
||||
|
|
@ -211,6 +220,29 @@ def parse_agent_litellm_params(value: object) -> Mapping[str, object]:
|
|||
return _EMPTY_LITELLM_PARAMS
|
||||
|
||||
|
||||
_KILL_SWITCH_ADAPTER: Final[TypeAdapter[AgentKillSwitchConfig | None]] = TypeAdapter(AgentKillSwitchConfig | None)
|
||||
|
||||
|
||||
def parse_agent_kill_switch(value: object) -> AgentKillSwitchConfig | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
if isinstance(value, str):
|
||||
return _KILL_SWITCH_ADAPTER.validate_json(value)
|
||||
return _KILL_SWITCH_ADAPTER.validate_python(value)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def serialize_agent_kill_switch(incoming: object, existing: object) -> str:
|
||||
"""prisma-client-py drops ``None`` from update data, so a cleared kill switch is stored as the JSON literal
|
||||
``null`` (read back as ``None``), the same convention ``memory_endpoints`` uses for ``Json?`` columns."""
|
||||
restored: Final = restore_kill_switch(
|
||||
_KILL_SWITCH_ADAPTER.validate_python(incoming), parse_agent_kill_switch(existing)
|
||||
)
|
||||
return safe_dumps(restored.model_dump() if restored is not None else None)
|
||||
|
||||
|
||||
_MISSING_AGENT_PARAM: Final = object()
|
||||
_RESTORE_AGENT_PARAMS_MAX_DEPTH: Final = 10
|
||||
|
||||
|
|
@ -293,6 +325,12 @@ def _patched_access_group_ids(agent: PatchAgentRequest) -> Mapping[str, object]:
|
|||
return MappingProxyType({"access_group_ids": tuple(dict.fromkeys(agent.get("access_group_ids") or ()))})
|
||||
|
||||
|
||||
def _patched_kill_switch(agent: PatchAgentRequest, existing: object) -> Mapping[str, object]:
|
||||
if "kill_switch" not in agent:
|
||||
return MappingProxyType({})
|
||||
return MappingProxyType({"kill_switch": serialize_agent_kill_switch(agent.get("kill_switch"), existing)})
|
||||
|
||||
|
||||
def _restore_redacted_litellm_params(
|
||||
incoming: Mapping[str, object],
|
||||
existing: Mapping[str, object],
|
||||
|
|
@ -531,6 +569,7 @@ class AgentRegistry:
|
|||
"agent_name": agent_name,
|
||||
"litellm_params": litellm_params,
|
||||
"agent_card_params": agent_card_params,
|
||||
"kill_switch": serialize_agent_kill_switch(agent.get("kill_switch"), None),
|
||||
"created_by": created_by,
|
||||
"updated_by": created_by,
|
||||
"created_at": datetime.now(timezone.utc),
|
||||
|
|
@ -613,7 +652,10 @@ class AgentRegistry:
|
|||
existing_agent: Final[Mapping[str, object]] = dict(existing_record)
|
||||
|
||||
augment_agent: Final = {**existing_agent, **agent}
|
||||
update_data: Final[dict[str, object]] = {**_patched_access_group_ids(agent)}
|
||||
update_data: Final[dict[str, object]] = {
|
||||
**_patched_access_group_ids(agent),
|
||||
**_patched_kill_switch(agent, existing_agent.get("kill_switch")),
|
||||
}
|
||||
if augment_agent.get("agent_name"):
|
||||
update_data["agent_name"] = augment_agent.get("agent_name")
|
||||
if "litellm_params" in agent:
|
||||
|
|
@ -716,6 +758,9 @@ class AgentRegistry:
|
|||
)
|
||||
extra_headers_val_u: Final = agent.get("extra_headers") or []
|
||||
access_group_ids_val_u: Final = tuple(dict.fromkeys(agent.get("access_group_ids") or ()))
|
||||
kill_switch_val_u: Final = serialize_agent_kill_switch(
|
||||
agent.get("kill_switch"), existing_row.kill_switch if existing_row is not None else None
|
||||
)
|
||||
|
||||
update_data: Final[dict[str, object]] = {
|
||||
"agent_name": agent_name,
|
||||
|
|
@ -723,6 +768,7 @@ class AgentRegistry:
|
|||
"agent_card_params": agent_card_params,
|
||||
"static_headers": static_headers_val_u,
|
||||
"extra_headers": extra_headers_val_u,
|
||||
"kill_switch": kill_switch_val_u,
|
||||
"access_group_ids": access_group_ids_val_u,
|
||||
"updated_by": updated_by,
|
||||
"updated_at": datetime.now(timezone.utc),
|
||||
|
|
|
|||
|
|
@ -33,6 +33,8 @@ from litellm.proxy.a2a.agent_card import (
|
|||
normalize_protocol_version,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.agent_registry import (
|
||||
AgentIdWhere,
|
||||
parse_agent_kill_switch,
|
||||
parse_agent_litellm_params,
|
||||
redact_sensitive_agent_litellm_params,
|
||||
)
|
||||
|
|
@ -45,6 +47,15 @@ from litellm.proxy.agent_endpoints.agent_search import (
|
|||
search_agents,
|
||||
)
|
||||
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import accessible_agents
|
||||
from litellm.proxy.agent_endpoints.kill_switch import (
|
||||
KillSwitchAuditLogWriter,
|
||||
KillSwitchHttpClient,
|
||||
build_kill_switch_audit_log,
|
||||
default_kill_switch_audit_log_writer,
|
||||
default_kill_switch_http_client,
|
||||
fire_kill_switch,
|
||||
redact_kill_switch,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user
|
||||
from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity
|
||||
|
|
@ -53,6 +64,8 @@ from litellm.types.agents import (
|
|||
AgentCard,
|
||||
AgentConfig,
|
||||
AgentKeySummary,
|
||||
AgentKillSwitchConfig,
|
||||
AgentKillSwitchResult,
|
||||
AgentMakePublicResponse,
|
||||
AgentResponse,
|
||||
MakeAgentsPublicRequest,
|
||||
|
|
@ -160,9 +173,10 @@ def _redact_sensitive_agent_fields(
|
|||
) -> list[AgentResponse]:
|
||||
"""
|
||||
Return copies of the given agents with credential-bearing litellm_params
|
||||
values replaced by a fixed marker (never returned to ANY caller,
|
||||
admin included) and, for non-admin callers, virtual-key and header
|
||||
fields stripped entirely. The original objects are not modified.
|
||||
values and kill-switch auth secrets replaced by a fixed marker (never
|
||||
returned to ANY caller, admin included) and, for non-admin callers,
|
||||
virtual-key, header and kill-switch fields stripped entirely. The original
|
||||
objects are not modified.
|
||||
"""
|
||||
redacted: Final[list[AgentResponse]] = []
|
||||
for agent in agents:
|
||||
|
|
@ -171,8 +185,10 @@ def _redact_sensitive_agent_fields(
|
|||
copy.static_headers = None
|
||||
copy.extra_headers = None
|
||||
copy.keys = None
|
||||
copy.kill_switch = None
|
||||
if copy.litellm_params:
|
||||
copy.litellm_params = _redact_agent_litellm_params_dict(copy.litellm_params)
|
||||
copy.kill_switch = redact_kill_switch(copy.kill_switch)
|
||||
redacted.append(copy)
|
||||
return redacted
|
||||
|
||||
|
|
@ -872,6 +888,74 @@ async def delete_agent(
|
|||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/agents/{agent_id}/kill_switch",
|
||||
tags=["[beta] A2A Agents"], # mutable-ok: fastapi types tags as list[str | Enum]
|
||||
dependencies=(Depends(user_api_key_auth),),
|
||||
response_model=AgentKillSwitchResult,
|
||||
)
|
||||
async def trigger_agent_kill_switch(
|
||||
agent_id: str,
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
http_client: Annotated[KillSwitchHttpClient, Depends(default_kill_switch_http_client)],
|
||||
audit_log_writer: Annotated[KillSwitchAuditLogWriter, Depends(default_kill_switch_audit_log_writer)],
|
||||
):
|
||||
"""
|
||||
Fire the agent's configured kill switch webhook. Proxy admin only.
|
||||
|
||||
LiteLLM only makes the configured HTTP call and reports what came back; it
|
||||
does not change the agent's state in LiteLLM. Returns 200 when the webhook
|
||||
answered 2xx, 502 with the same result body otherwise. Every attempt is
|
||||
written to the audit log as a `kill_switch_fired` row against the agent.
|
||||
|
||||
Example Request:
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000/kill_switch" \\
|
||||
-H "Authorization: Bearer <your_api_key>"
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import litellm_proxy_admin_name
|
||||
|
||||
await check_feature_access_for_user(user_api_key_dict, "agents")
|
||||
_check_agent_management_permission(user_api_key_dict)
|
||||
|
||||
resolved: Final = await _resolve_agent_kill_switch(agent_id)
|
||||
if resolved is None:
|
||||
raise HTTPException(status_code=404, detail=f"Agent with ID {agent_id} not found")
|
||||
resolved_agent_id, config = resolved
|
||||
if config is None:
|
||||
raise HTTPException(status_code=400, detail=f"Agent with ID {agent_id} has no kill_switch configured")
|
||||
|
||||
result: Final = await fire_kill_switch(agent_id=resolved_agent_id, config=config, http_client=http_client)
|
||||
await audit_log_writer(
|
||||
build_kill_switch_audit_log(
|
||||
result=result,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_proxy_admin_name=litellm_proxy_admin_name,
|
||||
)
|
||||
)
|
||||
if not result.succeeded:
|
||||
raise HTTPException(status_code=502, detail=result.model_dump())
|
||||
return result
|
||||
|
||||
|
||||
async def _resolve_agent_kill_switch(agent_id: str) -> tuple[str, AgentKillSwitchConfig | None] | None:
|
||||
"""The DB row wins over this replica's in-memory registry so a trigger never fires a webhook another
|
||||
replica has since changed; config.yaml agents have no row and fall back to the registry."""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is not None:
|
||||
where: Final[AgentIdWhere] = {"agent_id": agent_id}
|
||||
row: Final = await agents_table(prisma_client).find_unique(where=where)
|
||||
if row is not None:
|
||||
return row.agent_id, parse_agent_kill_switch(row.kill_switch)
|
||||
|
||||
agent: Final = AGENT_REGISTRY.get_agent_by_id(agent_id=agent_id)
|
||||
if agent is None:
|
||||
return None
|
||||
return agent.agent_id, agent.kill_switch
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/agents/{agent_id}/make_public",
|
||||
tags=["[beta] A2A Agents"],
|
||||
|
|
|
|||
239
litellm/proxy/agent_endpoints/kill_switch.py
Normal file
239
litellm/proxy/agent_endpoints/kill_switch.py
Normal file
|
|
@ -0,0 +1,239 @@
|
|||
from base64 import b64encode
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Protocol, TypeAlias
|
||||
|
||||
import httpx
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.constants import (
|
||||
AGENT_KILL_SWITCH_RESPONSE_BODY_MAX_CHARS,
|
||||
AGENT_KILL_SWITCH_TIMEOUT_SECONDS,
|
||||
REDACTED_BY_LITELM_STRING,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # its params arg is a bare dict in http_handler
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_AuditLogs, LitellmTableNames, UserAPIKeyAuth
|
||||
from litellm.proxy.management_helpers.audit_logs import create_audit_log_for_update, get_audit_log_changed_by
|
||||
from litellm.types.agents import (
|
||||
AgentKillSwitchApiKeyAuth,
|
||||
AgentKillSwitchAuth,
|
||||
AgentKillSwitchBasicAuth,
|
||||
AgentKillSwitchBearerAuth,
|
||||
AgentKillSwitchConfig,
|
||||
AgentKillSwitchResult,
|
||||
)
|
||||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
|
||||
|
||||
def _with_auth(config: AgentKillSwitchConfig, auth: AgentKillSwitchAuth) -> AgentKillSwitchConfig:
|
||||
return AgentKillSwitchConfig(
|
||||
url=config.url,
|
||||
method=config.method,
|
||||
headers=config.headers,
|
||||
query_params=config.query_params,
|
||||
body=config.body,
|
||||
auth=auth,
|
||||
)
|
||||
|
||||
|
||||
def redact_kill_switch(config: AgentKillSwitchConfig | None) -> AgentKillSwitchConfig | None:
|
||||
if config is None or config.auth is None:
|
||||
return config
|
||||
return _with_auth(config, _redact_auth(config.auth))
|
||||
|
||||
|
||||
def _redact_auth(auth: AgentKillSwitchAuth) -> AgentKillSwitchAuth:
|
||||
match auth:
|
||||
case AgentKillSwitchBearerAuth():
|
||||
return AgentKillSwitchBearerAuth(type="bearer", token=REDACTED_BY_LITELM_STRING)
|
||||
case AgentKillSwitchApiKeyAuth():
|
||||
return AgentKillSwitchApiKeyAuth(
|
||||
type="api_key", header_name=auth.header_name, api_key=REDACTED_BY_LITELM_STRING
|
||||
)
|
||||
case AgentKillSwitchBasicAuth():
|
||||
return AgentKillSwitchBasicAuth(type="basic", username=auth.username, password=REDACTED_BY_LITELM_STRING)
|
||||
case _:
|
||||
assert_never(auth)
|
||||
|
||||
|
||||
def restore_kill_switch(
|
||||
incoming: AgentKillSwitchConfig | None,
|
||||
existing: AgentKillSwitchConfig | None,
|
||||
) -> AgentKillSwitchConfig | None:
|
||||
"""Put the stored secret back behind an auth field echoed as the redaction
|
||||
marker; a marker with no stored secret of the same auth type becomes ""."""
|
||||
if incoming is None or incoming.auth is None:
|
||||
return incoming
|
||||
existing_auth: Final = existing.auth if existing is not None else None
|
||||
return _with_auth(incoming, _restore_auth(incoming.auth, existing_auth))
|
||||
|
||||
|
||||
def _restore_secret(incoming_value: str, existing_value: str | None) -> str:
|
||||
if incoming_value != REDACTED_BY_LITELM_STRING:
|
||||
return incoming_value
|
||||
return existing_value if existing_value is not None else ""
|
||||
|
||||
|
||||
def _restore_auth(incoming: AgentKillSwitchAuth, existing: AgentKillSwitchAuth | None) -> AgentKillSwitchAuth:
|
||||
match incoming:
|
||||
case AgentKillSwitchBearerAuth():
|
||||
stored_token: Final = existing.token if isinstance(existing, AgentKillSwitchBearerAuth) else None
|
||||
return AgentKillSwitchBearerAuth(type="bearer", token=_restore_secret(incoming.token, stored_token))
|
||||
case AgentKillSwitchApiKeyAuth():
|
||||
stored_key: Final = existing.api_key if isinstance(existing, AgentKillSwitchApiKeyAuth) else None
|
||||
return AgentKillSwitchApiKeyAuth(
|
||||
type="api_key",
|
||||
header_name=incoming.header_name,
|
||||
api_key=_restore_secret(incoming.api_key, stored_key),
|
||||
)
|
||||
case AgentKillSwitchBasicAuth():
|
||||
stored_password: Final = existing.password if isinstance(existing, AgentKillSwitchBasicAuth) else None
|
||||
return AgentKillSwitchBasicAuth(
|
||||
type="basic",
|
||||
username=incoming.username,
|
||||
password=_restore_secret(incoming.password, stored_password),
|
||||
)
|
||||
case _:
|
||||
assert_never(incoming)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class KillSwitchRequest:
|
||||
method: str
|
||||
url: str
|
||||
headers: Mapping[str, str]
|
||||
json_body: Mapping[str, object] | None
|
||||
|
||||
|
||||
def _auth_headers(auth: AgentKillSwitchAuth | None) -> Mapping[str, str]:
|
||||
match auth:
|
||||
case None:
|
||||
return MappingProxyType({})
|
||||
case AgentKillSwitchBearerAuth():
|
||||
return MappingProxyType({"Authorization": f"Bearer {auth.token}"})
|
||||
case AgentKillSwitchApiKeyAuth():
|
||||
return MappingProxyType({auth.header_name: auth.api_key})
|
||||
case AgentKillSwitchBasicAuth():
|
||||
credentials: Final = b64encode(f"{auth.username}:{auth.password}".encode()).decode()
|
||||
return MappingProxyType({"Authorization": f"Basic {credentials}"})
|
||||
case _:
|
||||
assert_never(auth)
|
||||
|
||||
|
||||
def build_kill_switch_request(config: AgentKillSwitchConfig) -> KillSwitchRequest:
|
||||
url: Final = httpx.URL(config.url).copy_merge_params(config.query_params)
|
||||
return KillSwitchRequest(
|
||||
method=config.method,
|
||||
url=str(url),
|
||||
headers=MappingProxyType({**config.headers, **_auth_headers(config.auth)}),
|
||||
json_body=config.body,
|
||||
)
|
||||
|
||||
|
||||
class KillSwitchHttpClient(Protocol):
|
||||
def build_request(
|
||||
self,
|
||||
method: str,
|
||||
url: str,
|
||||
*,
|
||||
headers: Mapping[str, str],
|
||||
json: Mapping[str, object] | None,
|
||||
timeout: float,
|
||||
) -> httpx.Request: ...
|
||||
|
||||
async def send(self, request: httpx.Request, *, stream: bool, follow_redirects: bool) -> httpx.Response: ...
|
||||
|
||||
|
||||
def default_kill_switch_http_client() -> KillSwitchHttpClient:
|
||||
return get_async_httpx_client(llm_provider=httpxSpecialProvider.AgentKillSwitch).client
|
||||
|
||||
|
||||
KillSwitchAuditLogWriter: TypeAlias = Callable[[LiteLLM_AuditLogs], Awaitable[None]] # mutable-ok: Callable params
|
||||
|
||||
|
||||
def default_kill_switch_audit_log_writer() -> KillSwitchAuditLogWriter:
|
||||
return create_audit_log_for_update
|
||||
|
||||
|
||||
def build_kill_switch_audit_log(
|
||||
*,
|
||||
result: AgentKillSwitchResult,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
litellm_proxy_admin_name: str | None,
|
||||
) -> LiteLLM_AuditLogs:
|
||||
return LiteLLM_AuditLogs(
|
||||
id=str(uuid.uuid4()),
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
changed_by=get_audit_log_changed_by(
|
||||
litellm_changed_by=None,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_proxy_admin_name=litellm_proxy_admin_name,
|
||||
),
|
||||
changed_by_api_key=user_api_key_dict.api_key,
|
||||
table_name=LitellmTableNames.AGENT_TABLE_NAME,
|
||||
object_id=result.agent_id,
|
||||
action="kill_switch_fired",
|
||||
updated_values=result.model_dump_json(exclude_none=True),
|
||||
)
|
||||
|
||||
|
||||
async def fire_kill_switch(
|
||||
*,
|
||||
agent_id: str,
|
||||
config: AgentKillSwitchConfig,
|
||||
http_client: KillSwitchHttpClient,
|
||||
timeout: float = AGENT_KILL_SWITCH_TIMEOUT_SECONDS,
|
||||
) -> AgentKillSwitchResult:
|
||||
request: Final = build_kill_switch_request(config)
|
||||
reported_url: Final = str(httpx.URL(request.url).copy_with(query=None))
|
||||
verbose_proxy_logger.info("Firing kill switch for agent %s: %s %s", agent_id, request.method, reported_url)
|
||||
try:
|
||||
response: Final = await http_client.send(
|
||||
http_client.build_request(
|
||||
request.method,
|
||||
request.url,
|
||||
headers=request.headers,
|
||||
json=request.json_body,
|
||||
timeout=timeout,
|
||||
),
|
||||
stream=True,
|
||||
follow_redirects=False,
|
||||
)
|
||||
body: Final = await _read_text_prefix(response, AGENT_KILL_SWITCH_RESPONSE_BODY_MAX_CHARS)
|
||||
except httpx.HTTPError as exc:
|
||||
verbose_proxy_logger.warning("Kill switch for agent %s failed: %s", agent_id, type(exc).__name__)
|
||||
return AgentKillSwitchResult(
|
||||
agent_id=agent_id,
|
||||
url=reported_url,
|
||||
method=config.method,
|
||||
error=type(exc).__name__,
|
||||
)
|
||||
return AgentKillSwitchResult(
|
||||
agent_id=agent_id,
|
||||
url=reported_url,
|
||||
method=config.method,
|
||||
status_code=response.status_code,
|
||||
response_body=body,
|
||||
)
|
||||
|
||||
|
||||
async def _read_text_prefix(response: httpx.Response, max_chars: int) -> str:
|
||||
try:
|
||||
return await _take_text(response.aiter_text(), max_chars)
|
||||
finally:
|
||||
await response.aclose()
|
||||
|
||||
|
||||
async def _take_text(chunks: AsyncIterator[str], max_chars: int) -> str:
|
||||
taken = "" # rebind-ok: running prefix of a stream that is abandoned once the cap is hit
|
||||
async for chunk in chunks:
|
||||
taken += chunk # rebind-ok: see above
|
||||
if len(taken) >= max_chars:
|
||||
break
|
||||
return taken[:max_chars]
|
||||
|
|
@ -293,6 +293,7 @@ from litellm.constants import (
|
|||
REALTIME_SESSION_FAILURE_LOGGED_KEY,
|
||||
REALTIME_SESSION_SUCCESS_LOGGED_KEY,
|
||||
ROUTER_SETTINGS_MANAGED_OUTSIDE_CONFIG,
|
||||
SPEND_CAPTURE_RATE_CHECK_JOB_ID,
|
||||
USER_SPEND_ALERTS_JOB_ID,
|
||||
WEEKLY_SPEND_REPORT_JOB_ID,
|
||||
)
|
||||
|
|
@ -758,6 +759,9 @@ from litellm.proxy.spend_tracking.budget_reservation import (
|
|||
from litellm.proxy.spend_tracking.daily_global_spend_rollup import (
|
||||
run_scheduled_daily_global_spend_reconcile,
|
||||
)
|
||||
from litellm.proxy.spend_tracking.spend_capture_rate import (
|
||||
run_scheduled_spend_capture_rate_check,
|
||||
)
|
||||
from litellm.proxy.spend_tracking.spend_counter_batch import (
|
||||
PendingSpendIncrement,
|
||||
active_spend_counter_batch,
|
||||
|
|
@ -856,6 +860,7 @@ from litellm.types.proxy.model_deprecation import (
|
|||
DEFAULT_DEPRECATION_WARN_DAYS,
|
||||
ModelDeprecationResponse,
|
||||
)
|
||||
from litellm.types.proxy.spend_capture_rate import SpendCaptureProvider, SpendCaptureRateCheckSettings
|
||||
from litellm.types.realtime import RealtimeQueryParams
|
||||
from litellm.types.router import (
|
||||
ClassifierPlugin,
|
||||
|
|
@ -5060,6 +5065,11 @@ def _bind_general_settings_store(settings: SettingsStore) -> None:
|
|||
general_settings = settings # pyright: ignore[reportAssignmentType] # legacy global accepts mappings
|
||||
|
||||
|
||||
def _current_general_settings() -> Mapping[str, object]:
|
||||
"""The live ``general_settings``, whichever object a config reload has bound since the caller was created."""
|
||||
return general_settings
|
||||
|
||||
|
||||
@lru_cache(maxsize=4096)
|
||||
def _log_ignored_cost_map_copy(model_id: str, fields: tuple[str, ...]) -> None:
|
||||
verbose_proxy_logger.warning(
|
||||
|
|
@ -10448,6 +10458,13 @@ class ProxyStartupEvent:
|
|||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
cls._initialize_spend_capture_rate_check_job(
|
||||
scheduler=scheduler,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
prisma_client=prisma_client,
|
||||
read_general_settings=_current_general_settings,
|
||||
)
|
||||
|
||||
### PTU DAILY ROLLUP ###
|
||||
from litellm.proxy.spend_tracking.ptu_feature_flag import (
|
||||
is_ptu_cost_attribution_enabled,
|
||||
|
|
@ -10822,6 +10839,64 @@ class ProxyStartupEvent:
|
|||
next_run_time=datetime.now(timezone.utc) + timedelta(minutes=2),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _initialize_spend_capture_rate_check_job(
|
||||
cls,
|
||||
scheduler: AsyncIOScheduler,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
prisma_client: PrismaClient,
|
||||
read_general_settings: Callable[[], Mapping[str, object]],
|
||||
) -> None:
|
||||
"""The job always runs and re-reads ``spend_capture_rate_check`` each run; an absent setting clears the gauge."""
|
||||
cls._spend_capture_rate_check_settings(read_general_settings())
|
||||
|
||||
async def alert(message: str) -> None:
|
||||
await proxy_logging_obj.alerting_handler(
|
||||
message=message,
|
||||
level="High",
|
||||
alert_type=AlertType.failed_tracking_spend,
|
||||
)
|
||||
|
||||
def publish(provider: str, capture_rate: float | None) -> None:
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
|
||||
for logger in litellm.logging_callback_manager.get_custom_loggers_for_type(callback_type=PrometheusLogger):
|
||||
if isinstance(logger, PrometheusLogger):
|
||||
logger.set_spend_capture_rate(api_provider=provider, capture_rate=capture_rate)
|
||||
|
||||
async def check() -> None:
|
||||
settings: Final = cls._spend_capture_rate_check_settings(read_general_settings())
|
||||
if settings is None:
|
||||
for provider in get_args(SpendCaptureProvider):
|
||||
publish(provider, None)
|
||||
return
|
||||
await run_scheduled_spend_capture_rate_check(
|
||||
prisma_client,
|
||||
settings,
|
||||
pod_lock_manager=proxy_logging_obj.db_spend_update_writer.pod_lock_manager,
|
||||
alert=alert,
|
||||
publish=publish,
|
||||
)
|
||||
|
||||
scheduler.add_job(
|
||||
check,
|
||||
"cron",
|
||||
hour=1,
|
||||
minute=15,
|
||||
timezone="UTC",
|
||||
id=SPEND_CAPTURE_RATE_CHECK_JOB_ID,
|
||||
replace_existing=True,
|
||||
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
|
||||
next_run_time=datetime.now(timezone.utc) + timedelta(minutes=2),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _spend_capture_rate_check_settings(
|
||||
general_settings: Mapping[str, object],
|
||||
) -> SpendCaptureRateCheckSettings | None:
|
||||
raw_settings: Final = general_settings.get("spend_capture_rate_check")
|
||||
return None if raw_settings is None else SpendCaptureRateCheckSettings.model_validate(raw_settings)
|
||||
|
||||
@classmethod
|
||||
async def _initialize_slack_alerting_jobs(
|
||||
cls,
|
||||
|
|
@ -13766,7 +13841,7 @@ async def transform_request(request: TransformRequestBody):
|
|||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail={"error": str(e)})
|
||||
|
||||
return return_raw_request(endpoint=request.call_type, kwargs=request.request_body)
|
||||
return await asyncio.to_thread(return_raw_request, request.call_type, request.request_body)
|
||||
|
||||
|
||||
async def _check_if_model_is_user_added(
|
||||
|
|
|
|||
|
|
@ -72,6 +72,7 @@ model LiteLLM_AgentsTable {
|
|||
agent_card_params Json
|
||||
static_headers Json? @default("{}")
|
||||
extra_headers String[] @default([])
|
||||
kill_switch Json?
|
||||
agent_access_groups String[] @default([])
|
||||
access_group_ids String[] @default([])
|
||||
object_permission_id String?
|
||||
|
|
|
|||
300
litellm/proxy/spend_tracking/spend_capture_rate.py
Normal file
300
litellm/proxy/spend_tracking/spend_capture_rate.py
Normal file
|
|
@ -0,0 +1,300 @@
|
|||
"""Compare the spend LiteLLM captured for a provider against what that provider billed for the same UTC days.
|
||||
|
||||
LiteLLM's side is ``LiteLLM_DailyUserSpend``, summed over the ``custom_llm_provider`` values that land on the
|
||||
provider's bill. The provider's side is its billing API, read with the customer's own billing credential
|
||||
(OpenAI: the organization costs endpoint and an admin key in ``OPENAI_ADMIN_KEY``).
|
||||
"""
|
||||
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from datetime import date, datetime, timedelta, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, TypeAlias
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import (
|
||||
SPEND_CAPTURE_RATE_CHECK_JOB_ID,
|
||||
SPEND_CAPTURE_RATE_CHECK_LOCK_TTL_SECONDS,
|
||||
SPEND_CAPTURE_RATE_DOCS_URL,
|
||||
)
|
||||
from litellm.llms.openai.organization_costs import (
|
||||
OPENAI_ADMIN_KEY_ENV_VAR,
|
||||
BillingHttpGet,
|
||||
OpenAICostsRequestFailed,
|
||||
fetch_openai_daily_costs,
|
||||
provider_billing_get,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.proxy.spend_capture_rate import (
|
||||
CaptureRateDay,
|
||||
CaptureRateReport,
|
||||
SpendCaptureProvider,
|
||||
SpendCaptureRateCheckSettings,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
OPENAI_BILLED_LITELLM_PROVIDERS: Final = ("openai", "text-completion-openai")
|
||||
|
||||
CaptureRatePublisher: TypeAlias = Callable[[SpendCaptureProvider, float | None], None] # mutable-ok: Callable params
|
||||
|
||||
_CAPTURED_SPEND_BY_DAY_SQL: Final = """
|
||||
SELECT date, COALESCE(SUM(spend), 0)::float AS spend
|
||||
FROM "LiteLLM_DailyUserSpend"
|
||||
WHERE date >= $1 AND date <= $2 AND custom_llm_provider = ANY($3::text[])
|
||||
GROUP BY date
|
||||
"""
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ProviderBillingCredentialMissing:
|
||||
provider: SpendCaptureProvider
|
||||
env_var: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ProviderBillingRequestFailed:
|
||||
provider: SpendCaptureProvider
|
||||
detail: str
|
||||
|
||||
|
||||
ProviderBillingFailure: TypeAlias = ProviderBillingCredentialMissing | ProviderBillingRequestFailed
|
||||
CheckResult: TypeAlias = CaptureRateReport | ProviderBillingFailure
|
||||
|
||||
|
||||
class _CapturedSpendRow(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
date: str
|
||||
spend: float
|
||||
|
||||
|
||||
_CAPTURED_SPEND_ROWS: Final = TypeAdapter(tuple[_CapturedSpendRow, ...])
|
||||
|
||||
|
||||
async def captured_spend_by_day(
|
||||
prisma_client: "PrismaClient",
|
||||
*,
|
||||
litellm_providers: Sequence[str],
|
||||
start_date: date,
|
||||
end_date: date,
|
||||
) -> Mapping[str, float]:
|
||||
"""LiteLLM's tracked spend per UTC day (ISO date) for the given ``custom_llm_provider`` values."""
|
||||
rows: Final = await prisma_client.db.query_raw(
|
||||
_CAPTURED_SPEND_BY_DAY_SQL, start_date.isoformat(), end_date.isoformat(), tuple(litellm_providers)
|
||||
)
|
||||
return MappingProxyType({row.date: row.spend for row in _CAPTURED_SPEND_ROWS.validate_python(rows)})
|
||||
|
||||
|
||||
def _ratio(captured: float, billed: float) -> float | None:
|
||||
return None if billed <= 0 else captured / billed
|
||||
|
||||
|
||||
def _days(start_date: date, end_date: date) -> tuple[date, ...]:
|
||||
return tuple(start_date + timedelta(days=offset) for offset in range((end_date - start_date).days + 1))
|
||||
|
||||
|
||||
def compute_capture_rate(
|
||||
*,
|
||||
provider: SpendCaptureProvider,
|
||||
start_date: date,
|
||||
end_date: date,
|
||||
captured_by_day: Mapping[str, float],
|
||||
billed_by_day: Mapping[str, float],
|
||||
threshold: float,
|
||||
) -> CaptureRateReport:
|
||||
days: Final = tuple(
|
||||
CaptureRateDay(
|
||||
date=day.isoformat(),
|
||||
captured_spend=captured_by_day.get(day.isoformat(), 0.0),
|
||||
provider_spend=billed_by_day.get(day.isoformat(), 0.0),
|
||||
capture_rate=_ratio(captured_by_day.get(day.isoformat(), 0.0), billed_by_day.get(day.isoformat(), 0.0)),
|
||||
)
|
||||
for day in _days(start_date, end_date)
|
||||
)
|
||||
captured: Final = sum(day.captured_spend for day in days)
|
||||
billed: Final = sum(day.provider_spend for day in days)
|
||||
rate: Final = _ratio(captured, billed)
|
||||
return CaptureRateReport(
|
||||
provider=provider,
|
||||
start_date=start_date.isoformat(),
|
||||
end_date=end_date.isoformat(),
|
||||
captured_spend=captured,
|
||||
provider_spend=billed,
|
||||
capture_rate=rate,
|
||||
threshold=threshold,
|
||||
below_threshold=rate is not None and rate < threshold,
|
||||
days=days,
|
||||
)
|
||||
|
||||
|
||||
async def capture_rate_report(
|
||||
prisma_client: "PrismaClient",
|
||||
*,
|
||||
provider: SpendCaptureProvider,
|
||||
start_date: date,
|
||||
end_date: date,
|
||||
threshold: float,
|
||||
openai_project_ids: Sequence[str] = (),
|
||||
http_get: BillingHttpGet = provider_billing_get,
|
||||
) -> CheckResult:
|
||||
match provider:
|
||||
case "openai":
|
||||
admin_key: Final = get_secret_str(OPENAI_ADMIN_KEY_ENV_VAR)
|
||||
if admin_key is None:
|
||||
return ProviderBillingCredentialMissing(provider, OPENAI_ADMIN_KEY_ENV_VAR)
|
||||
billed: Final = await fetch_openai_daily_costs(
|
||||
start_date, end_date, admin_key=admin_key, project_ids=openai_project_ids, http_get=http_get
|
||||
)
|
||||
if isinstance(billed, OpenAICostsRequestFailed):
|
||||
return ProviderBillingRequestFailed(provider, billed.detail)
|
||||
captured: Final = await captured_spend_by_day(
|
||||
prisma_client,
|
||||
litellm_providers=OPENAI_BILLED_LITELLM_PROVIDERS,
|
||||
start_date=start_date,
|
||||
end_date=end_date,
|
||||
)
|
||||
return compute_capture_rate(
|
||||
provider=provider,
|
||||
start_date=start_date,
|
||||
end_date=end_date,
|
||||
captured_by_day=captured,
|
||||
billed_by_day=billed,
|
||||
threshold=threshold,
|
||||
)
|
||||
case _:
|
||||
assert_never(provider)
|
||||
|
||||
|
||||
def alert_message(result: CheckResult) -> str | None:
|
||||
"""The alert a check outcome warrants, or ``None`` when the capture rate is healthy."""
|
||||
match result:
|
||||
case ProviderBillingCredentialMissing(provider=provider, env_var=env_var):
|
||||
return (
|
||||
f"Spend capture-rate check: {env_var} is not set, so the {provider} bill cannot be read. "
|
||||
f"Set it or remove general_settings.spend_capture_rate_check. {SPEND_CAPTURE_RATE_DOCS_URL}"
|
||||
)
|
||||
case ProviderBillingRequestFailed(provider=provider, detail=detail):
|
||||
return f"Spend capture-rate check: could not read the {provider} bill ({detail}). {SPEND_CAPTURE_RATE_DOCS_URL}"
|
||||
case CaptureRateReport():
|
||||
if not result.below_threshold or result.capture_rate is None:
|
||||
return None
|
||||
return (
|
||||
f"Spend capture rate for {result.provider} is {result.capture_rate:.1%}, under the "
|
||||
f"{result.threshold:.0%} threshold: LiteLLM captured ${result.captured_spend:,.2f} of the "
|
||||
f"${result.provider_spend:,.2f} {result.provider} bill for {result.start_date} to {result.end_date}. "
|
||||
f"Requests reach {result.provider} outside LiteLLM or cost tracking is dropping spend. "
|
||||
f"{SPEND_CAPTURE_RATE_DOCS_URL}"
|
||||
)
|
||||
case _:
|
||||
assert_never(result)
|
||||
|
||||
|
||||
def _published_rate(result: CheckResult) -> float | None:
|
||||
"""The gauge value: the rate, or ``None`` (NaN on the gauge) when this window produced no rate."""
|
||||
return result.capture_rate if isinstance(result, CaptureRateReport) else None
|
||||
|
||||
|
||||
async def _check_every_provider(
|
||||
prisma_client: "PrismaClient",
|
||||
settings: SpendCaptureRateCheckSettings,
|
||||
*,
|
||||
publish: CaptureRatePublisher,
|
||||
today: date | None,
|
||||
http_get: BillingHttpGet,
|
||||
) -> tuple[CheckResult, ...]:
|
||||
"""Check every configured provider over the closed days before ``today`` and publish each outcome."""
|
||||
end_date: Final = (today or datetime.now(timezone.utc).date()) - timedelta(days=1)
|
||||
start_date: Final = end_date - timedelta(days=settings.lookback_days - 1)
|
||||
results: Final = tuple(
|
||||
[
|
||||
await capture_rate_report(
|
||||
prisma_client,
|
||||
provider=provider,
|
||||
start_date=start_date,
|
||||
end_date=end_date,
|
||||
threshold=settings.threshold,
|
||||
openai_project_ids=settings.openai_project_ids,
|
||||
http_get=http_get,
|
||||
)
|
||||
for provider in settings.providers
|
||||
]
|
||||
)
|
||||
for result in results:
|
||||
publish(result.provider, _published_rate(result))
|
||||
verbose_proxy_logger.info("Spend capture-rate check: %s", result)
|
||||
return results
|
||||
|
||||
|
||||
def _alert_messages(results: Sequence[CheckResult]) -> tuple[str, ...]:
|
||||
return tuple(message for message in map(alert_message, results) if message is not None)
|
||||
|
||||
|
||||
async def run_spend_capture_rate_check(
|
||||
prisma_client: "PrismaClient",
|
||||
settings: SpendCaptureRateCheckSettings,
|
||||
*,
|
||||
alert: Callable[[str], Awaitable[None]],
|
||||
publish: CaptureRatePublisher,
|
||||
today: date | None = None,
|
||||
http_get: BillingHttpGet = provider_billing_get,
|
||||
) -> tuple[CheckResult, ...]:
|
||||
"""Check every configured provider, publish each rate, and alert on every outcome that warrants one."""
|
||||
results: Final = await _check_every_provider(
|
||||
prisma_client, settings, publish=publish, today=today, http_get=http_get
|
||||
)
|
||||
for message in _alert_messages(results):
|
||||
await alert(message)
|
||||
return results
|
||||
|
||||
|
||||
async def run_scheduled_spend_capture_rate_check(
|
||||
prisma_client: "PrismaClient",
|
||||
settings: SpendCaptureRateCheckSettings,
|
||||
*,
|
||||
pod_lock_manager: "PodLockManager | None",
|
||||
alert: Callable[[str], Awaitable[None]],
|
||||
publish: CaptureRatePublisher,
|
||||
today: date | None = None,
|
||||
http_get: BillingHttpGet = provider_billing_get,
|
||||
) -> tuple[CheckResult, ...]:
|
||||
"""Every worker publishes its own gauge; the first replica whose finished check has an alert claims the window."""
|
||||
results: Final = await _check_every_provider(
|
||||
prisma_client, settings, publish=publish, today=today, http_get=http_get
|
||||
)
|
||||
messages: Final = _alert_messages(results)
|
||||
if not messages:
|
||||
return results
|
||||
if not await _claims_alert_window(pod_lock_manager):
|
||||
verbose_proxy_logger.info("Spend capture-rate check: another pod alerted this window")
|
||||
return results
|
||||
for message in messages:
|
||||
await alert(message)
|
||||
return results
|
||||
|
||||
|
||||
async def _claims_alert_window(pod_lock_manager: "PodLockManager | None") -> bool:
|
||||
"""The lock is left to expire, so every replica firing within its TTL of the winner stays quiet."""
|
||||
redis_cache: Final = None if pod_lock_manager is None else pod_lock_manager.redis_cache
|
||||
if pod_lock_manager is None or redis_cache is None:
|
||||
return True
|
||||
acquired: Final = await pod_lock_manager.acquire_lock(
|
||||
cronjob_id=SPEND_CAPTURE_RATE_CHECK_JOB_ID, ttl=SPEND_CAPTURE_RATE_CHECK_LOCK_TTL_SECONDS
|
||||
)
|
||||
return acquired or not await _lock_is_held(pod_lock_manager, redis_cache)
|
||||
|
||||
|
||||
async def _lock_is_held(pod_lock_manager: "PodLockManager", redis_cache: "RedisCache") -> bool:
|
||||
try:
|
||||
return bool(
|
||||
await redis_cache.async_get_cache(pod_lock_manager.get_redis_lock_key(SPEND_CAPTURE_RATE_CHECK_JOB_ID))
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 # an unreadable lock must not silence the alert
|
||||
verbose_proxy_logger.warning("Spend capture-rate check: could not read the lock: %s", exc)
|
||||
return False
|
||||
|
|
@ -24,7 +24,7 @@ from typing import (
|
|||
import fastapi
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
|
||||
from pydantic import TypeAdapter
|
||||
from typing_extensions import ReadOnly
|
||||
from typing_extensions import ReadOnly, assert_never
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -32,11 +32,17 @@ from litellm.constants import (
|
|||
EMPTY_MAPPING,
|
||||
LITELLM_TRUNCATED_PAYLOAD_FIELD,
|
||||
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
|
||||
SPEND_CAPTURE_RATE_MAX_RANGE_DAYS,
|
||||
)
|
||||
from litellm.litellm_core_utils.classifier_logging import classifier_audit_fields, classifier_input_snapshot
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy._types import ProviderBudgetResponse, ProviderBudgetResponseObject
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.spend_tracking.spend_capture_rate import (
|
||||
ProviderBillingCredentialMissing,
|
||||
ProviderBillingRequestFailed,
|
||||
capture_rate_report,
|
||||
)
|
||||
|
||||
# NOTE: Avoid module-level import from common_utils: proxy_server imports this
|
||||
# module while common_utils may pull proxy_server during init, which can leave
|
||||
|
|
@ -52,6 +58,7 @@ from litellm.repositories.team_repository import TeamRepository
|
|||
from litellm.repositories.verification_token_repository import (
|
||||
VerificationTokenRepository,
|
||||
)
|
||||
from litellm.types.proxy.spend_capture_rate import CaptureRateReport, SpendCaptureProvider
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from prisma import models as prisma_models
|
||||
|
|
@ -1183,6 +1190,84 @@ async def get_global_activity_exceptions(
|
|||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/spend/capture_rate",
|
||||
tags=["Budget & Spend Tracking"], # mutable-ok: FastAPI tags kwarg is list-typed
|
||||
dependencies=(Depends(user_api_key_auth),),
|
||||
response_model=CaptureRateReport,
|
||||
)
|
||||
async def get_spend_capture_rate(
|
||||
start_date: Annotated[date, fastapi.Query(description="First UTC day of the range, YYYY-MM-DD")],
|
||||
end_date: Annotated[date, fastapi.Query(description="Last UTC day of the range, YYYY-MM-DD, inclusive")],
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
provider: Annotated[
|
||||
SpendCaptureProvider,
|
||||
fastapi.Query(description="Provider whose bill to compare against; needs OPENAI_ADMIN_KEY set on the proxy"),
|
||||
] = "openai",
|
||||
threshold: Annotated[
|
||||
float, fastapi.Query(gt=0, le=1, description="Ratio under which the report flags below_threshold")
|
||||
] = 0.9,
|
||||
project_ids: Annotated[
|
||||
list[str] | None,
|
||||
fastapi.Query(
|
||||
description=(
|
||||
"Scope the OpenAI bill to these project ids; omit to compare against the whole organization. Captured "
|
||||
"spend is never scoped, so pass every project LiteLLM's OpenAI keys belong to"
|
||||
)
|
||||
),
|
||||
] = None,
|
||||
) -> CaptureRateReport:
|
||||
"""
|
||||
Compare the spend LiteLLM captured for a provider against that provider's own bill, per UTC day.
|
||||
|
||||
Admin only. Reads the provider's billing API with the billing credential set on the proxy
|
||||
(OpenAI: `OPENAI_ADMIN_KEY`) and sums `LiteLLM_DailyUserSpend` for the same days.
|
||||
|
||||
Example:
|
||||
```
|
||||
curl -H "Authorization: Bearer sk-1234" \
|
||||
"http://localhost:4000/spend/capture_rate?provider=openai&start_date=2026-09-17&end_date=2026-09-23"
|
||||
```
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if not _is_admin_view_safe(user_api_key_dict):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Only proxy admins can read the capture rate")
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail=CommonProxyErrors.db_not_connected_error.value
|
||||
)
|
||||
if end_date < start_date:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="end_date must not be before start_date")
|
||||
if (end_date - start_date).days >= SPEND_CAPTURE_RATE_MAX_RANGE_DAYS:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Date range too large; maximum is {SPEND_CAPTURE_RATE_MAX_RANGE_DAYS} days",
|
||||
)
|
||||
result: Final = await capture_rate_report(
|
||||
prisma_client,
|
||||
provider=provider,
|
||||
start_date=start_date,
|
||||
end_date=end_date,
|
||||
threshold=threshold,
|
||||
openai_project_ids=tuple(project_ids or ()),
|
||||
)
|
||||
match result:
|
||||
case ProviderBillingCredentialMissing(env_var=env_var):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail=f"{env_var} is not set on the proxy, so the {provider} bill cannot be read",
|
||||
)
|
||||
case ProviderBillingRequestFailed(detail=detail):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY, detail=f"Could not read the {provider} bill: {detail}"
|
||||
)
|
||||
case CaptureRateReport():
|
||||
return result
|
||||
case _:
|
||||
assert_never(result)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/global/spend/provider",
|
||||
tags=["Budget & Spend Tracking"],
|
||||
|
|
|
|||
|
|
@ -1,8 +1,9 @@
|
|||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
from typing import TYPE_CHECKING, Annotated, Any, Final, Literal, TypeAlias
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, PrivateAttr, StrictInt
|
||||
from pydantic import BaseModel, ConfigDict, Field, PrivateAttr, StrictInt, field_validator
|
||||
from typing_extensions import ReadOnly, Required, TypedDict
|
||||
|
||||
from litellm.types.llms.base import LiteLLMPydanticObjectBase
|
||||
|
|
@ -178,6 +179,74 @@ class AgentObjectPermission(TypedDict, total=False):
|
|||
agents: list[str] | None
|
||||
|
||||
|
||||
class AgentKillSwitchBearerAuth(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="forbid")
|
||||
|
||||
type: Literal["bearer"]
|
||||
token: str
|
||||
|
||||
|
||||
class AgentKillSwitchApiKeyAuth(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="forbid")
|
||||
|
||||
type: Literal["api_key"]
|
||||
header_name: str = "x-api-key"
|
||||
api_key: str
|
||||
|
||||
|
||||
class AgentKillSwitchBasicAuth(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="forbid")
|
||||
|
||||
type: Literal["basic"]
|
||||
username: str
|
||||
password: str
|
||||
|
||||
|
||||
AgentKillSwitchAuth: TypeAlias = Annotated[
|
||||
AgentKillSwitchBearerAuth | AgentKillSwitchApiKeyAuth | AgentKillSwitchBasicAuth,
|
||||
Field(discriminator="type"),
|
||||
]
|
||||
|
||||
AgentKillSwitchMethod: TypeAlias = Literal["POST", "PUT", "PATCH", "DELETE", "GET"]
|
||||
|
||||
|
||||
class AgentKillSwitchConfig(BaseModel):
|
||||
"""Webhook an admin fires to shut an agent down out of band. LiteLLM only
|
||||
makes the call; whatever the endpoint does with it is the agent's business."""
|
||||
|
||||
model_config = ConfigDict(frozen=True, extra="forbid")
|
||||
|
||||
url: str
|
||||
method: AgentKillSwitchMethod = "POST"
|
||||
headers: Mapping[str, str] = Field(default_factory=dict)
|
||||
query_params: Mapping[str, str] = Field(default_factory=dict)
|
||||
body: Mapping[str, object] | None = None
|
||||
auth: AgentKillSwitchAuth | None = None
|
||||
|
||||
@field_validator("url")
|
||||
@classmethod
|
||||
def _require_absolute_http_url(cls, value: str) -> str:
|
||||
parts: Final = urlsplit(value)
|
||||
if parts.scheme not in ("http", "https") or not parts.netloc:
|
||||
raise ValueError("kill_switch.url must be an absolute http(s) URL")
|
||||
return value
|
||||
|
||||
|
||||
class AgentKillSwitchResult(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
agent_id: str
|
||||
url: str
|
||||
method: AgentKillSwitchMethod
|
||||
status_code: int | None = None
|
||||
response_body: str | None = None
|
||||
error: str | None = None
|
||||
|
||||
@property
|
||||
def succeeded(self) -> bool:
|
||||
return self.status_code is not None and 200 <= self.status_code < 300
|
||||
|
||||
|
||||
class AgentConfig(TypedDict, total=False):
|
||||
agent_name: Required[str]
|
||||
agent_card_params: Required[AgentCard]
|
||||
|
|
@ -190,6 +259,7 @@ class AgentConfig(TypedDict, total=False):
|
|||
static_headers: dict[str, str] | None
|
||||
extra_headers: list[str] | None
|
||||
access_group_ids: ReadOnly[Sequence[str] | None]
|
||||
kill_switch: ReadOnly[AgentKillSwitchConfig | None]
|
||||
|
||||
|
||||
class PatchAgentRequest(TypedDict, total=False):
|
||||
|
|
@ -204,6 +274,7 @@ class PatchAgentRequest(TypedDict, total=False):
|
|||
static_headers: dict[str, str] | None
|
||||
extra_headers: list[str] | None
|
||||
access_group_ids: ReadOnly[Sequence[str] | None]
|
||||
kill_switch: ReadOnly[AgentKillSwitchConfig | None]
|
||||
|
||||
|
||||
AGENT_CALLER_USER_ID_HEADER: Final = "x-litellm-user-id"
|
||||
|
|
@ -243,6 +314,7 @@ class AgentResponse(BaseModel):
|
|||
static_headers: dict[str, str] | None = None
|
||||
extra_headers: list[str] | None = None
|
||||
access_group_ids: Sequence[str] | None = None
|
||||
kill_switch: AgentKillSwitchConfig | None = None
|
||||
keys: list[AgentKeySummary] | None = None
|
||||
search_score: float | None = None
|
||||
created_at: datetime | None = None
|
||||
|
|
|
|||
|
|
@ -281,6 +281,7 @@ DEFINED_PROMETHEUS_METRICS = Literal[
|
|||
"litellm_guardrail_errors_total",
|
||||
"litellm_guardrail_requests_total",
|
||||
"litellm_zero_cost_requests_total",
|
||||
"litellm_spend_capture_rate",
|
||||
# Cache metrics
|
||||
"litellm_cache_hits_metric",
|
||||
"litellm_cache_misses_metric",
|
||||
|
|
@ -600,6 +601,8 @@ class PrometheusMetricLabels:
|
|||
ZERO_COST_REASON_LABEL,
|
||||
)
|
||||
|
||||
litellm_spend_capture_rate = (UserAPIKeyLabelNames.API_PROVIDER.value,)
|
||||
|
||||
litellm_input_tokens_metric = [
|
||||
UserAPIKeyLabelNames.END_USER.value,
|
||||
UserAPIKeyLabelNames.API_KEY_HASH.value,
|
||||
|
|
|
|||
|
|
@ -24,8 +24,10 @@ class httpxSpecialProvider(str, Enum):
|
|||
Search = "search"
|
||||
MCP = "mcp"
|
||||
RAG = "rag"
|
||||
ProviderBilling = "provider_billing"
|
||||
A2AProvider = "a2a_provider"
|
||||
AgentHealthCheck = "agent_health_check"
|
||||
AgentKillSwitch = "agent_kill_switch"
|
||||
A2A = "a2a"
|
||||
PromptManagement = "prompt_management"
|
||||
UI = "ui"
|
||||
|
|
|
|||
54
litellm/types/proxy/spend_capture_rate.py
Normal file
54
litellm/types/proxy/spend_capture_rate.py
Normal file
|
|
@ -0,0 +1,54 @@
|
|||
"""The captured-spend to provider-bill ratio: the share of a provider's bill that went through LiteLLM and was priced.
|
||||
|
||||
``capture_rate = captured_spend / provider_spend`` over the same UTC days. 1.0 means LiteLLM saw and priced every
|
||||
dollar the provider billed, lower means traffic reaches the provider outside LiteLLM or cost tracking drops spend,
|
||||
higher means LiteLLM prices above the bill. ``None`` means the provider billed nothing, so there is no ratio.
|
||||
"""
|
||||
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
from litellm.constants import SPEND_CAPTURE_RATE_MAX_RANGE_DAYS
|
||||
|
||||
SpendCaptureProvider = Literal["openai"]
|
||||
|
||||
|
||||
class SpendCaptureRateCheckSettings(BaseModel):
|
||||
"""``general_settings.spend_capture_rate_check``: the daily check of captured spend against the provider bill."""
|
||||
|
||||
model_config = ConfigDict(frozen=True, extra="forbid")
|
||||
|
||||
providers: tuple[SpendCaptureProvider, ...] = Field(("openai",), min_length=1)
|
||||
threshold: float = Field(0.9, gt=0, le=1)
|
||||
lookback_days: int = Field(7, ge=1, le=SPEND_CAPTURE_RATE_MAX_RANGE_DAYS)
|
||||
openai_project_ids: tuple[str, ...] = Field(
|
||||
(),
|
||||
description=(
|
||||
"Scope the OpenAI bill to these project ids; empty compares against the whole organization. Captured "
|
||||
"spend is never scoped, so list every project LiteLLM's OpenAI keys belong to"
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class CaptureRateDay(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
date: str
|
||||
captured_spend: float
|
||||
provider_spend: float
|
||||
capture_rate: float | None
|
||||
|
||||
|
||||
class CaptureRateReport(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
provider: SpendCaptureProvider
|
||||
start_date: str
|
||||
end_date: str
|
||||
captured_spend: float
|
||||
provider_spend: float
|
||||
capture_rate: float | None
|
||||
threshold: float
|
||||
below_threshold: bool
|
||||
days: tuple[CaptureRateDay, ...]
|
||||
|
|
@ -10273,7 +10273,7 @@ def return_raw_request(endpoint: CallTypes, kwargs: dict) -> RawRequestTypedDict
|
|||
"""
|
||||
from datetime import datetime
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging, RawRequestCaptured
|
||||
|
||||
litellm_logging_obj: Final = Logging(
|
||||
model="gpt-3.5-turbo",
|
||||
|
|
@ -10284,6 +10284,7 @@ def return_raw_request(endpoint: CallTypes, kwargs: dict) -> RawRequestTypedDict
|
|||
start_time=datetime.now(),
|
||||
function_id="1234",
|
||||
log_raw_request_response=True,
|
||||
raw_request_only=True,
|
||||
)
|
||||
|
||||
llm_api_endpoint: Final = getattr(litellm, endpoint.value)
|
||||
|
|
@ -10294,7 +10295,11 @@ def return_raw_request(endpoint: CallTypes, kwargs: dict) -> RawRequestTypedDict
|
|||
llm_api_endpoint(
|
||||
**kwargs,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
api_key="my-fake-api-key", # 👈 ensure the request fails
|
||||
api_key="my-fake-api-key",
|
||||
)
|
||||
except RawRequestCaptured:
|
||||
received_exception = (
|
||||
"raw request was not captured before the provider call; check the proxy logs for the pre_call error"
|
||||
)
|
||||
except Exception as e:
|
||||
received_exception = str(e)
|
||||
|
|
|
|||
|
|
@ -26207,7 +26207,6 @@
|
|||
},
|
||||
"gemini-2.5-flash-image": {
|
||||
"deprecation_date": "2027-03-15",
|
||||
"cache_read_input_token_cost": 3e-08,
|
||||
"input_cost_per_audio_token": 1e-06,
|
||||
"input_cost_per_token": 3e-07,
|
||||
"input_cost_per_token_batches": 1.5e-07,
|
||||
|
|
@ -26312,10 +26311,14 @@
|
|||
"gemini-3-pro-image-preview": {
|
||||
"input_cost_per_image": 0.0011,
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_priority": 3.6e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07,
|
||||
"cache_read_input_token_cost_batches": 1e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_priority": 3.6e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_above_200k_tokens_priority": 7.2e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 65536,
|
||||
|
|
@ -26325,7 +26328,9 @@
|
|||
"output_cost_per_image": 0.134,
|
||||
"output_cost_per_image_token": 0.00012,
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_priority": 2.16e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_above_200k_tokens_priority": 3.24e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"source": "https://ai.google.dev/gemini-api/docs/pricing",
|
||||
"supported_endpoints": [
|
||||
|
|
@ -27848,6 +27853,7 @@
|
|||
"gemini-embedding-001": {
|
||||
"deprecation_date": "2028-05-20",
|
||||
"input_cost_per_token": 1.5e-07,
|
||||
"input_cost_per_token_batches": 1.2e-07,
|
||||
"litellm_provider": "vertex_ai-embedding-models",
|
||||
"max_input_tokens": 2048,
|
||||
"max_tokens": 2048,
|
||||
|
|
@ -49410,7 +49416,6 @@
|
|||
},
|
||||
"vertex_ai/gemini-2.5-flash-image": {
|
||||
"deprecation_date": "2027-03-15",
|
||||
"cache_read_input_token_cost": 3e-08,
|
||||
"input_cost_per_audio_token": 1e-06,
|
||||
"input_cost_per_token": 3e-07,
|
||||
"input_cost_per_token_batches": 1.5e-07,
|
||||
|
|
@ -49492,10 +49497,14 @@
|
|||
"vertex_ai/gemini-3-pro-image-preview": {
|
||||
"input_cost_per_image": 0.0011,
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_priority": 3.6e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_priority": 7.2e-07,
|
||||
"cache_read_input_token_cost_batches": 1e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_priority": 3.6e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 4e-06,
|
||||
"input_cost_per_token_above_200k_tokens_priority": 7.2e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "vertex_ai-language-models",
|
||||
"max_input_tokens": 65536,
|
||||
|
|
@ -49505,7 +49514,9 @@
|
|||
"output_cost_per_image": 0.134,
|
||||
"output_cost_per_image_token": 0.00012,
|
||||
"output_cost_per_token": 1.2e-05,
|
||||
"output_cost_per_token_priority": 2.16e-05,
|
||||
"output_cost_per_token_above_200k_tokens": 1.8e-05,
|
||||
"output_cost_per_token_above_200k_tokens_priority": 3.24e-05,
|
||||
"output_cost_per_token_batches": 6e-06,
|
||||
"supports_reasoning": false,
|
||||
"source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image"
|
||||
|
|
@ -69473,6 +69484,7 @@
|
|||
"input_cost_per_audio_token": 3e-06,
|
||||
"input_cost_per_image_token": 1e-06,
|
||||
"input_cost_per_token": 7.5e-07,
|
||||
"input_cost_per_video_token": 1e-06,
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 131072,
|
||||
"max_output_tokens": 65536,
|
||||
|
|
@ -69493,6 +69505,7 @@
|
|||
"input_cost_per_audio_token": 3e-06,
|
||||
"input_cost_per_image_token": 1e-06,
|
||||
"input_cost_per_token": 7.5e-07,
|
||||
"input_cost_per_video_token": 1e-06,
|
||||
"litellm_provider": "gemini",
|
||||
"max_input_tokens": 131072,
|
||||
"max_output_tokens": 65536,
|
||||
|
|
|
|||
|
|
@ -72,6 +72,7 @@ model LiteLLM_AgentsTable {
|
|||
agent_card_params Json
|
||||
static_headers Json? @default("{}")
|
||||
extra_headers String[] @default([])
|
||||
kill_switch Json?
|
||||
agent_access_groups String[] @default([])
|
||||
access_group_ids String[] @default([])
|
||||
object_permission_id String?
|
||||
|
|
|
|||
|
|
@ -98,7 +98,7 @@ Each suite provides its own `client` fixture (see `llm_translation/passthrough_c
|
|||
|
||||
Request and response bodies are typed pydantic models in `models.py`; only the fields a test reads are modelled, and nothing passes raw dicts. Outcomes come back as a `Result[R]` tagged union (`Success`, `NetworkError`, `UnauthorizedError`, `RateLimitedError`, `ValidationError`, `UnknownApiError`). Handle them with `match`, or call `unwrap(...)` when a non-success should fail the test. The harness hard-fails and never skips: a test marked `e2e` fails when no proxy answers its liveness probe, and once a request reaches the proxy any wrong behavior is likewise a hard failure, so a missing proxy turns the run red instead of being mistaken for a pass
|
||||
|
||||
Mark live tests with `@pytest.mark.e2e` (on the class or the module). Add `@pytest.mark.quiet_stack` to a test that measures the proxy itself (RSS, latency): the shared stack lock in `stack_lock.py` then runs it while no other test on the host is hitting the stack, marked or not, so the reading depends only on the test's own traffic. Use `scoped_key` for a fresh all-models key that auto-deletes, `resources` when you need to create and tear down more than a key, and `unique_marker()` from `e2e_config` to keep prompts, tags, and customer ids from colliding across concurrent runs and the shared response cache
|
||||
Mark live tests with `@pytest.mark.e2e` (on the class or the module). Coverage of the harness itself carries no marker and runs whether or not a proxy is up. Add `@pytest.mark.quiet_stack` to a test that measures the proxy itself (RSS, latency): the shared stack lock in `stack_lock.py` then runs it while no other test on the host is hitting the stack, marked or not, so the reading depends only on the test's own traffic. Use `scoped_key` for a fresh all-models key that auto-deletes, `resources` when you need to create and tear down more than a key, and `unique_marker()` from `e2e_config` to keep prompts, tags, and customer ids from colliding across concurrent runs and the shared response cache
|
||||
|
||||
## Record and replay fixtures
|
||||
|
||||
|
|
@ -251,7 +251,7 @@ other.<area>.<case>.<assertion>
|
|||
```
|
||||
|
||||
## Hard Rules
|
||||
- no unit tests of any kind under `tests/e2e`. a product feature is proven end to end against a live proxy, never with a unit test, and the harness itself is not unit-tested here either. no monkeypatching or mock tests. if a contributor asks you to write an end to end test, do NOT stage a unit test with it; if you find a product gap, call it out in the PR description
|
||||
- no unit tests of a product feature under `tests/e2e`, and no mock tests or monkeypatching of code anywhere in it: a product feature is proven end to end against a live proxy, never with a unit test. if a contributor asks you to write an end to end test, do NOT stage a unit test with it; if you find a product gap, call it out in the PR description. the harness's own plumbing is the one exception: the markerless tests in the root-level `test_*.py` files, `coverage_registry/test_collector.py`, `guardrails/test_guardrails_client.py`, the `claude_code/_*_unit_tests/` trees, and the `load/` aggregation tests carry no `e2e` marker, run without a proxy, and take their inputs as arguments or env vars (setting an env var through pytest's `monkeypatch` fixture is fine, patching a function, class, or module is not), and no coverage-registry or compat-matrix cell rests on them. judge a change inside one of them by that standard, not as a misplaced product test
|
||||
|
||||
- use model management endpoints to create new models for a test. this could be in a conftest / inline for each test. ask the user what they want.
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,61 @@
|
|||
import math
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from prometheus_client import REGISTRY
|
||||
from prometheus_client.samples import Sample
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
|
||||
METRIC: Final = "litellm_spend_capture_rate"
|
||||
|
||||
|
||||
def _clear_prometheus_registry() -> None:
|
||||
for collector in list(REGISTRY._collector_to_names.keys()):
|
||||
try:
|
||||
REGISTRY.unregister(collector)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _samples(metric_name: str) -> list[Sample]:
|
||||
return [sample for metric in REGISTRY.collect() for sample in metric.samples if sample.name == metric_name]
|
||||
|
||||
|
||||
def test_capture_rate_gauge_holds_the_latest_rate_per_provider_and_nan_when_there_is_none() -> None:
|
||||
_clear_prometheus_registry()
|
||||
try:
|
||||
logger: Final = PrometheusLogger()
|
||||
assert _samples(METRIC) == []
|
||||
|
||||
logger.set_spend_capture_rate(api_provider="openai", capture_rate=0.87)
|
||||
logger.set_spend_capture_rate(api_provider="openai", capture_rate=0.91)
|
||||
|
||||
samples: Final = _samples(METRIC)
|
||||
assert [(sample.labels, sample.value) for sample in samples] == [({"api_provider": "openai"}, 0.91)]
|
||||
|
||||
logger.set_spend_capture_rate(api_provider="openai", capture_rate=None)
|
||||
|
||||
(unavailable,) = _samples(METRIC)
|
||||
assert unavailable.labels == {"api_provider": "openai"} and math.isnan(unavailable.value)
|
||||
finally:
|
||||
_clear_prometheus_registry()
|
||||
|
||||
|
||||
def test_capture_rate_gauge_still_records_when_api_provider_is_an_excluded_label(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(litellm, "prometheus_exclude_labels", ["api_provider"])
|
||||
_clear_prometheus_registry()
|
||||
try:
|
||||
logger: Final = PrometheusLogger()
|
||||
|
||||
logger.set_spend_capture_rate(api_provider="openai", capture_rate=0.42)
|
||||
|
||||
assert [(sample.labels, sample.value) for sample in _samples(METRIC)] == [({}, 0.42)]
|
||||
|
||||
logger.set_spend_capture_rate(api_provider="openai", capture_rate=None)
|
||||
|
||||
(unavailable,) = _samples(METRIC)
|
||||
assert unavailable.labels == {} and math.isnan(unavailable.value)
|
||||
finally:
|
||||
_clear_prometheus_registry()
|
||||
112
tests/test_litellm/llms/openai/test_organization_costs.py
Normal file
112
tests/test_litellm/llms/openai/test_organization_costs.py
Normal file
|
|
@ -0,0 +1,112 @@
|
|||
from collections.abc import Mapping
|
||||
from datetime import date, datetime, timezone
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.constants import OPENAI_ORGANIZATION_COSTS_URL
|
||||
from litellm.llms.openai.organization_costs import OpenAICostsRequestFailed, fetch_openai_daily_costs
|
||||
|
||||
_ADMIN_KEY: Final = "sk-admin-test"
|
||||
|
||||
|
||||
def _utc_midnight(day: str) -> int:
|
||||
return int(datetime.fromisoformat(day).replace(tzinfo=timezone.utc).timestamp())
|
||||
|
||||
|
||||
def _bucket(day: str, *amounts: float) -> dict[str, object]:
|
||||
return {
|
||||
"object": "bucket",
|
||||
"start_time": _utc_midnight(day),
|
||||
"end_time": _utc_midnight(day) + 86400,
|
||||
"results": [
|
||||
{"object": "organization.costs.result", "amount": {"value": amount, "currency": "usd"}, "line_item": None}
|
||||
for amount in amounts
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
class _FakeCostsApi:
|
||||
"""Serves ``pages`` in order and records every request it saw."""
|
||||
|
||||
def __init__(self, *pages: dict[str, object] | httpx.Response | Exception) -> None:
|
||||
self._pages = list(pages)
|
||||
self.calls: list[tuple[str, Mapping[str, object], Mapping[str, str]]] = []
|
||||
|
||||
async def __call__(self, url: str, params: Mapping[str, object], headers: Mapping[str, str]) -> httpx.Response:
|
||||
self.calls.append((url, dict(params), dict(headers)))
|
||||
page = self._pages.pop(0)
|
||||
if isinstance(page, Exception):
|
||||
raise page
|
||||
if isinstance(page, httpx.Response):
|
||||
return page
|
||||
return httpx.Response(200, json=page)
|
||||
|
||||
|
||||
def _page(*buckets: dict[str, object], next_page: str | None = None) -> dict[str, object]:
|
||||
return {"object": "page", "data": list(buckets), "has_more": next_page is not None, "next_page": next_page}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_costs_are_summed_per_utc_day_across_pages_and_line_items():
|
||||
api = _FakeCostsApi(
|
||||
_page(_bucket("2026-09-20", 10.0, 2.5), _bucket("2026-09-21", 4.0), next_page="page_2"),
|
||||
_page(_bucket("2026-09-22", 1.0)),
|
||||
)
|
||||
|
||||
billed = await fetch_openai_daily_costs(date(2026, 9, 20), date(2026, 9, 22), admin_key=_ADMIN_KEY, http_get=api)
|
||||
|
||||
assert dict(billed) == {"2026-09-20": 12.5, "2026-09-21": 4.0, "2026-09-22": 1.0}
|
||||
first, second = api.calls
|
||||
assert first[0] == OPENAI_ORGANIZATION_COSTS_URL
|
||||
assert first[2] == {"Authorization": f"Bearer {_ADMIN_KEY}"}
|
||||
assert first[1]["start_time"] == _utc_midnight("2026-09-20")
|
||||
assert first[1]["end_time"] == _utc_midnight("2026-09-23")
|
||||
assert first[1]["bucket_width"] == "1d"
|
||||
assert "page" not in first[1]
|
||||
assert "project_ids[]" not in first[1]
|
||||
assert second[1] == {**first[1], "page": "page_2"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openai_costs_are_scoped_to_the_configured_projects():
|
||||
api = _FakeCostsApi(_page())
|
||||
|
||||
billed = await fetch_openai_daily_costs(
|
||||
date(2026, 9, 20), date(2026, 9, 20), admin_key=_ADMIN_KEY, project_ids=("proj_a", "proj_b"), http_get=api
|
||||
)
|
||||
|
||||
assert dict(billed) == {}
|
||||
assert api.calls[0][1]["project_ids[]"] == ("proj_a", "proj_b")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"page, detail_fragment",
|
||||
[
|
||||
(httpx.Response(401, json={"error": {"message": "Incorrect API key provided"}}), "HTTP 401"),
|
||||
(httpx.ConnectError("connection refused"), "request failed"),
|
||||
({"object": "page", "data": [{"start_time": "not-a-timestamp"}]}, "unexpected response shape"),
|
||||
],
|
||||
)
|
||||
async def test_an_unreadable_openai_bill_is_a_request_failure_not_an_exception(page, detail_fragment):
|
||||
api = _FakeCostsApi(page)
|
||||
|
||||
billed = await fetch_openai_daily_costs(date(2026, 9, 20), date(2026, 9, 20), admin_key=_ADMIN_KEY, http_get=api)
|
||||
|
||||
assert isinstance(billed, OpenAICostsRequestFailed)
|
||||
assert detail_fragment in billed.detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_failure_on_a_later_page_fails_the_whole_read():
|
||||
api = _FakeCostsApi(
|
||||
_page(_bucket("2026-09-20", 10.0), next_page="page_2"),
|
||||
httpx.Response(429, json={"error": {"message": "rate limited"}}),
|
||||
)
|
||||
|
||||
billed = await fetch_openai_daily_costs(date(2026, 9, 20), date(2026, 9, 21), admin_key=_ADMIN_KEY, http_get=api)
|
||||
|
||||
assert isinstance(billed, OpenAICostsRequestFailed)
|
||||
assert "HTTP 429" in billed.detail
|
||||
|
|
@ -2,6 +2,7 @@ import contextlib
|
|||
import json
|
||||
import os
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Literal
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -4285,6 +4286,105 @@ class TestAgentMCPPermissions:
|
|||
):
|
||||
assert await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth=agent_key) == []
|
||||
|
||||
@staticmethod
|
||||
def _tool_grants(grants: dict[str, dict[str, list[str]]], keyed_by: Literal["team_id", "user_id"]) -> AsyncMock:
|
||||
"""Object permissions keyed by the ``team_id`` or ``user_id`` being asked about; anyone else has none."""
|
||||
|
||||
async def by_principal(user_api_key_auth: UserAPIKeyAuth | None = None) -> LiteLLM_ObjectPermissionTable | None:
|
||||
assert user_api_key_auth is not None
|
||||
principal = (user_api_key_auth.team_id if keyed_by == "team_id" else user_api_key_auth.user_id) or ""
|
||||
tools = grants.get(principal)
|
||||
if tools is None:
|
||||
return None
|
||||
return LiteLLM_ObjectPermissionTable(object_permission_id=f"perm-{principal}", mcp_tool_permissions=tools)
|
||||
|
||||
return AsyncMock(side_effect=by_principal)
|
||||
|
||||
@contextlib.contextmanager
|
||||
def _caller_tool_levels(
|
||||
self, team_grants: dict[str, dict[str, list[str]]], user_grants: dict[str, dict[str, list[str]]]
|
||||
):
|
||||
with (
|
||||
patch.object( # test-quality-ok: the level loaders read proxy_server globals with no injection seam
|
||||
MCPRequestHandler, "_get_key_object_permission", return_value=None
|
||||
),
|
||||
patch.object( # test-quality-ok: same seam, keyed by which team is being asked about
|
||||
MCPRequestHandler, "_get_team_object_permission", self._tool_grants(team_grants, keyed_by="team_id")
|
||||
),
|
||||
patch.object( # test-quality-ok: same seam, keyed by which user is being asked about
|
||||
MCPRequestHandler, "_get_user_object_permission", self._tool_grants(user_grants, keyed_by="user_id")
|
||||
),
|
||||
patch.object( # test-quality-ok: agent object_permission lookup hits the DB, not under test here
|
||||
MCPRequestHandler, "_get_agent_object_permission", AsyncMock(return_value=None)
|
||||
),
|
||||
):
|
||||
yield
|
||||
|
||||
async def test_agent_key_acting_for_a_user_only_sees_the_tools_that_user_may_call(self):
|
||||
"""The agent's key may call every tool on server-a, the invoking team grants two of them and the
|
||||
invoking user only one, so on that user's behalf the agent sees exactly that one tool."""
|
||||
agent_key = self._agent_key_acting_for(user_id="alice", team_id="callers")
|
||||
|
||||
with self._caller_tool_levels(
|
||||
team_grants={"callers": {"server-a": ["read_wiki_structure", "ask_wiki_question"]}},
|
||||
user_grants={"alice": {"server-a": ["read_wiki_structure", "read_wiki_contents"]}},
|
||||
):
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("server-a", agent_key) == [
|
||||
"read_wiki_structure"
|
||||
]
|
||||
|
||||
async def test_agent_key_acting_for_a_user_without_tool_grants_keeps_its_own_tools(self):
|
||||
agent_key = self._agent_key_acting_for(user_id="alice", team_id="callers")
|
||||
|
||||
with self._caller_tool_levels(team_grants={}, user_grants={}):
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("server-a", agent_key) is None
|
||||
|
||||
async def test_agent_key_acting_for_a_team_is_capped_at_that_teams_tools_on_the_server(self):
|
||||
agent_key = self._agent_key_acting_for(user_id="alice", team_id="callers")
|
||||
|
||||
with self._caller_tool_levels(
|
||||
team_grants={"callers": {"server-a": ["ask_wiki_question"], "server-b": ["other"]}}, user_grants={}
|
||||
):
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("server-a", agent_key) == ["ask_wiki_question"]
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("server-c", agent_key) is None
|
||||
|
||||
async def test_agent_key_not_acting_for_anyone_ignores_the_caller_tool_ceiling(self):
|
||||
agent_key = UserAPIKeyAuth(api_key="agent-key", user_id="agent-owner", team_id="agent-team", agent_id="agent-1")
|
||||
|
||||
with self._caller_tool_levels(
|
||||
team_grants={"callers": {"server-a": ["ask_wiki_question"]}},
|
||||
user_grants={"alice": {"server-a": ["read_wiki_structure"]}},
|
||||
):
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("server-a", agent_key) is None
|
||||
|
||||
async def test_agent_key_acting_for_a_caller_whose_team_is_unreadable_gets_no_tools(self):
|
||||
agent_key = self._agent_key_acting_for(user_id="alice", team_id="callers")
|
||||
|
||||
async def only_the_callers_team_is_unreadable(
|
||||
user_api_key_auth: UserAPIKeyAuth | None = None,
|
||||
) -> LiteLLM_ObjectPermissionTable | None:
|
||||
if user_api_key_auth is not None and user_api_key_auth.team_id == "callers":
|
||||
raise RuntimeError("db down")
|
||||
return None
|
||||
|
||||
with (
|
||||
patch.object( # test-quality-ok: the level loaders read proxy_server globals with no injection seam
|
||||
MCPRequestHandler, "_get_key_object_permission", return_value=None
|
||||
),
|
||||
patch.object( # test-quality-ok: same seam; the agent's own team resolves, the caller's team does not
|
||||
MCPRequestHandler,
|
||||
"_get_team_object_permission",
|
||||
AsyncMock(side_effect=only_the_callers_team_is_unreadable),
|
||||
),
|
||||
patch.object( # test-quality-ok: same seam
|
||||
MCPRequestHandler, "_get_user_object_permission", AsyncMock(return_value=None)
|
||||
),
|
||||
patch.object( # test-quality-ok: agent object_permission lookup hits the DB, not under test here
|
||||
MCPRequestHandler, "_get_agent_object_permission", AsyncMock(return_value=None)
|
||||
),
|
||||
):
|
||||
assert await MCPRequestHandler.get_allowed_tools_for_server("server-a", agent_key) == []
|
||||
|
||||
async def test_get_allowed_mcp_servers_agent_intersection(self):
|
||||
"""Key/team allow [server_1, server_2]; agent allows [server_1]. Result = [server_1]."""
|
||||
user_api_key_auth = UserAPIKeyAuth(
|
||||
|
|
@ -4355,8 +4455,7 @@ class TestAgentMCPPermissions:
|
|||
assert result == frozenset({"ag-server-id"})
|
||||
assert asked == ["agent-ag"]
|
||||
assert (
|
||||
await MCPRequestHandler._get_agent_access_group_server_ceiling(UserAPIKeyAuth(api_key="k"), resolve)
|
||||
is None
|
||||
await MCPRequestHandler._get_agent_access_group_server_ceiling(UserAPIKeyAuth(api_key="k"), resolve) is None
|
||||
)
|
||||
assert asked == ["agent-ag"]
|
||||
|
||||
|
|
@ -4493,7 +4592,9 @@ class TestAgentMCPPermissions:
|
|||
stack.enter_context(patcher)
|
||||
stack.enter_context(
|
||||
patch.object( # test-quality-ok: key resolution has its own tests; pin its grants here
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_key", AsyncMock(return_value=["server-a", "server-b"])
|
||||
MCPRequestHandler,
|
||||
"_get_allowed_mcp_servers_for_key",
|
||||
AsyncMock(return_value=["server-a", "server-b"]),
|
||||
)
|
||||
)
|
||||
stack.enter_context(
|
||||
|
|
@ -4519,7 +4620,9 @@ class TestAgentMCPPermissions:
|
|||
await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth)
|
||||
stack.enter_context(
|
||||
patch.object( # test-quality-ok: key resolution has its own tests; pin its grants here
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_key", AsyncMock(return_value=["server-a", "server-b"])
|
||||
MCPRequestHandler,
|
||||
"_get_allowed_mcp_servers_for_key",
|
||||
AsyncMock(return_value=["server-a", "server-b"]),
|
||||
)
|
||||
)
|
||||
stack.enter_context(
|
||||
|
|
@ -4543,9 +4646,15 @@ class TestAgentMCPPermissions:
|
|||
with contextlib.ExitStack() as stack:
|
||||
for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager):
|
||||
stack.enter_context(patcher)
|
||||
server_a_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server("server-a", user_api_key_auth)
|
||||
server_b_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server("server-b", user_api_key_auth)
|
||||
server_c_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server("server-c", user_api_key_auth)
|
||||
server_a_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server(
|
||||
"server-a", user_api_key_auth
|
||||
)
|
||||
server_b_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server(
|
||||
"server-b", user_api_key_auth
|
||||
)
|
||||
server_c_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server(
|
||||
"server-c", user_api_key_auth
|
||||
)
|
||||
|
||||
assert sorted(server_a_tools) == ["tool_direct", "tool_via_toolset"]
|
||||
assert server_b_tools == ["tool_b"]
|
||||
|
|
|
|||
|
|
@ -453,7 +453,7 @@ async def test_update_agent_in_db_raises_when_row_deleted_mid_update():
|
|||
registry: Final = AgentRegistry()
|
||||
mock_prisma: Final = MagicMock()
|
||||
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
|
||||
return_value=SimpleNamespace(litellm_params={}, object_permission_id=None)
|
||||
return_value=SimpleNamespace(litellm_params={}, object_permission_id=None, kill_switch=None)
|
||||
)
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_prisma.db.litellm_agentstable.update = AsyncMock(return_value=None)
|
||||
|
|
@ -742,6 +742,7 @@ async def test_update_agent_in_db_preserves_secret_when_echoed_back_redacted():
|
|||
"model": "bedrock/agentcore/my-agent",
|
||||
},
|
||||
object_permission_id=None,
|
||||
kill_switch=None,
|
||||
)
|
||||
)
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
|
|
@ -791,6 +792,7 @@ async def test_update_agent_in_db_preserves_secret_when_key_omitted_entirely():
|
|||
return_value=SimpleNamespace(
|
||||
litellm_params={"aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY},
|
||||
object_permission_id=None,
|
||||
kill_switch=None,
|
||||
)
|
||||
)
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
|
|
@ -838,6 +840,7 @@ async def test_update_agent_in_db_preserves_secret_nested_under_a_non_sensitive_
|
|||
}
|
||||
},
|
||||
object_permission_id=None,
|
||||
kill_switch=None,
|
||||
)
|
||||
)
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
|
|
@ -887,6 +890,7 @@ async def test_update_agent_in_db_clears_secret_on_explicit_empty_value():
|
|||
return_value=SimpleNamespace(
|
||||
litellm_params={"aws_secret_access_key": SENTINEL_AWS_SECRET_ACCESS_KEY},
|
||||
object_permission_id=None,
|
||||
kill_switch=None,
|
||||
)
|
||||
)
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
|
|
@ -1126,7 +1130,9 @@ async def test_update_agent_in_db_always_writes_access_group_ids(body_access_gro
|
|||
registry: Final = AgentRegistry()
|
||||
mock_prisma: Final = MagicMock()
|
||||
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
|
||||
return_value=SimpleNamespace(litellm_params={}, object_permission_id=None, access_group_ids=["ag-1"])
|
||||
return_value=SimpleNamespace(
|
||||
litellm_params={}, object_permission_id=None, kill_switch=None, access_group_ids=["ag-1"]
|
||||
)
|
||||
)
|
||||
mock_prisma.replica_db = mock_prisma.db
|
||||
mock_update = AsyncMock(return_value=_agent_row_mock(expected))
|
||||
|
|
@ -1143,3 +1149,155 @@ async def test_update_agent_in_db_always_writes_access_group_ids(body_access_gro
|
|||
)
|
||||
|
||||
assert tuple(mock_update.call_args.kwargs["data"]["access_group_ids"]) == tuple(expected)
|
||||
|
||||
|
||||
_KILL_SWITCH: Final = {
|
||||
"url": "https://ops.example.com/kill",
|
||||
"method": "POST",
|
||||
"headers": {"X-Env": "prod"},
|
||||
"query_params": {"reason": "manual"},
|
||||
"body": {"action": "stop"},
|
||||
"auth": {"type": "bearer", "token": "tok-real"},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_agent_to_db_stores_kill_switch_json_and_a_json_null_when_unset():
|
||||
registry: Final = AgentRegistry()
|
||||
mock_prisma: Final = MagicMock()
|
||||
mock_create = AsyncMock(return_value=_agent_row_mock([]))
|
||||
mock_prisma.db.litellm_agentstable.create = mock_create
|
||||
|
||||
await registry.add_agent_to_db(
|
||||
agent={
|
||||
"agent_name": "Test Agent",
|
||||
"agent_card_params": _sample_agent_card_params(),
|
||||
"kill_switch": _KILL_SWITCH,
|
||||
},
|
||||
prisma_client=mock_prisma,
|
||||
created_by="test-user",
|
||||
)
|
||||
assert json.loads(mock_create.call_args.kwargs["data"]["kill_switch"]) == _KILL_SWITCH
|
||||
|
||||
await registry.add_agent_to_db(
|
||||
agent={"agent_name": "Plain Agent", "agent_card_params": _sample_agent_card_params()},
|
||||
prisma_client=mock_prisma,
|
||||
created_by="test-user",
|
||||
)
|
||||
assert mock_create.call_args.kwargs["data"]["kill_switch"] == json.dumps(None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_agent_to_db_rejects_a_kill_switch_with_a_non_http_url():
|
||||
registry: Final = AgentRegistry()
|
||||
mock_prisma: Final = MagicMock()
|
||||
mock_prisma.db.litellm_agentstable.create = AsyncMock(return_value=_agent_row_mock([]))
|
||||
|
||||
with pytest.raises(Exception, match="absolute http"):
|
||||
await registry.add_agent_to_db(
|
||||
agent={
|
||||
"agent_name": "Test Agent",
|
||||
"agent_card_params": _sample_agent_card_params(),
|
||||
"kill_switch": {**_KILL_SWITCH, "url": "ops.example.com/kill"},
|
||||
},
|
||||
prisma_client=mock_prisma,
|
||||
created_by="test-user",
|
||||
)
|
||||
mock_prisma.db.litellm_agentstable.create.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_patch_agent_in_db_keeps_kill_switch_when_omitted_and_clears_it_on_null():
|
||||
registry: Final = AgentRegistry()
|
||||
mock_prisma: Final = MagicMock()
|
||||
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
|
||||
return_value={
|
||||
"agent_id": "agent-123",
|
||||
"agent_name": "Old",
|
||||
"litellm_params": {},
|
||||
"object_permission_id": None,
|
||||
"kill_switch": _KILL_SWITCH,
|
||||
}
|
||||
)
|
||||
mock_update = AsyncMock(return_value=_agent_row_mock([]))
|
||||
mock_prisma.db.litellm_agentstable.update = mock_update
|
||||
|
||||
await registry.patch_agent_in_db(
|
||||
agent_id="agent-123", agent={"agent_name": "New"}, prisma_client=mock_prisma, updated_by="u"
|
||||
)
|
||||
assert "kill_switch" not in mock_update.call_args.kwargs["data"]
|
||||
|
||||
await registry.patch_agent_in_db(
|
||||
agent_id="agent-123", agent={"kill_switch": None}, prisma_client=mock_prisma, updated_by="u"
|
||||
)
|
||||
assert mock_update.call_args.kwargs["data"]["kill_switch"] == json.dumps(None), (
|
||||
"prisma-client-py silently drops None, so the clear must be written as the JSON literal null"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_patch_agent_in_db_restores_the_stored_kill_switch_secret_behind_the_marker():
|
||||
registry: Final = AgentRegistry()
|
||||
mock_prisma: Final = MagicMock()
|
||||
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
|
||||
return_value={
|
||||
"agent_id": "agent-123",
|
||||
"agent_name": "A",
|
||||
"litellm_params": {},
|
||||
"object_permission_id": None,
|
||||
"kill_switch": _KILL_SWITCH,
|
||||
}
|
||||
)
|
||||
mock_update = AsyncMock(return_value=_agent_row_mock([]))
|
||||
mock_prisma.db.litellm_agentstable.update = mock_update
|
||||
|
||||
await registry.patch_agent_in_db(
|
||||
agent_id="agent-123",
|
||||
agent={
|
||||
"kill_switch": {
|
||||
**_KILL_SWITCH,
|
||||
"url": "https://ops.example.com/v2/kill",
|
||||
"auth": {"type": "bearer", "token": REDACTED_BY_LITELM_STRING},
|
||||
}
|
||||
},
|
||||
prisma_client=mock_prisma,
|
||||
updated_by="u",
|
||||
)
|
||||
|
||||
assert json.loads(mock_update.call_args.kwargs["data"]["kill_switch"]) == {
|
||||
**_KILL_SWITCH,
|
||||
"url": "https://ops.example.com/v2/kill",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_agent_in_db_clears_kill_switch_when_omitted_and_restores_secret_when_echoed():
|
||||
registry: Final = AgentRegistry()
|
||||
mock_prisma: Final = MagicMock()
|
||||
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(
|
||||
return_value=SimpleNamespace(litellm_params={}, object_permission_id=None, kill_switch=json.dumps(_KILL_SWITCH))
|
||||
)
|
||||
mock_update = AsyncMock(return_value=_agent_row_mock([]))
|
||||
mock_prisma.db.litellm_agentstable.update = mock_update
|
||||
base: Final = {"agent_name": "Test Agent", "agent_card_params": _sample_agent_card_params(), "litellm_params": {}}
|
||||
|
||||
await registry.update_agent_in_db(agent_id="agent-123", agent=base, prisma_client=mock_prisma, updated_by="u")
|
||||
assert mock_update.call_args.kwargs["data"]["kill_switch"] == json.dumps(None)
|
||||
|
||||
echoed: Final = {**_KILL_SWITCH, "auth": {"type": "bearer", "token": REDACTED_BY_LITELM_STRING}}
|
||||
await registry.update_agent_in_db(
|
||||
agent_id="agent-123", agent={**base, "kill_switch": echoed}, prisma_client=mock_prisma, updated_by="u"
|
||||
)
|
||||
assert json.loads(mock_update.call_args.kwargs["data"]["kill_switch"]) == _KILL_SWITCH
|
||||
|
||||
|
||||
def test_load_agents_from_config_exposes_a_typed_kill_switch():
|
||||
registry: Final = AgentRegistry()
|
||||
|
||||
registry.load_agents_from_config(
|
||||
[{"agent_name": "cfg-agent", "agent_card_params": _sample_agent_card_params(), "kill_switch": _KILL_SWITCH}]
|
||||
)
|
||||
|
||||
(agent,) = registry.get_agent_list()
|
||||
assert agent.kill_switch is not None
|
||||
assert agent.kill_switch.model_dump() == _KILL_SWITCH
|
||||
|
|
|
|||
|
|
@ -1,13 +1,15 @@
|
|||
import json
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from litellm.constants import REDACTED_BY_LITELM_STRING
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy._types import LiteLLM_AuditLogs, LitellmTableNames, LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.agent_endpoints import endpoints as agent_endpoints
|
||||
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import (
|
||||
RestrictedAgentAccess,
|
||||
|
|
@ -1147,3 +1149,222 @@ def test_make_agent_public_rejects_an_agent_published_only_in_the_db(monkeypatch
|
|||
|
||||
assert duplicate.status_code == 400
|
||||
assert "already in public agent groups" in duplicate.json()["detail"]
|
||||
|
||||
|
||||
_KILL_SWITCH: Final = {
|
||||
"url": "https://ops.example.com/kill",
|
||||
"method": "POST",
|
||||
"headers": {"X-Env": "prod"},
|
||||
"query_params": {"reason": "manual"},
|
||||
"body": {"action": "stop"},
|
||||
"auth": {"type": "bearer", "token": "tok-real"},
|
||||
}
|
||||
|
||||
|
||||
def _agent_with_kill_switch() -> AgentResponse:
|
||||
return AgentResponse(
|
||||
agent_id="agent-123",
|
||||
agent_name="Test Agent",
|
||||
agent_card_params=_sample_agent_card_params(),
|
||||
litellm_params={},
|
||||
kill_switch=_KILL_SWITCH,
|
||||
)
|
||||
|
||||
|
||||
class _FakeKillSwitchClient:
|
||||
def __init__(self, response: httpx.Response) -> None:
|
||||
self.calls: list[tuple[str, str, dict[str, str], object, float]] = [] # mutable-ok: test double records calls
|
||||
self._response: Final = response
|
||||
|
||||
def build_request(self, method: str, url: str, *, headers, json, timeout: float) -> httpx.Request:
|
||||
self.calls.append((method, url, dict(headers), json, timeout))
|
||||
return httpx.Request(method, url, headers=dict(headers), json=json)
|
||||
|
||||
async def send(self, request: httpx.Request, *, stream: bool, follow_redirects: bool) -> httpx.Response:
|
||||
return self._response
|
||||
|
||||
|
||||
class _AuditLogRecorder:
|
||||
def __init__(self) -> None:
|
||||
self.rows: list[LiteLLM_AuditLogs] = [] # mutable-ok: test double records writes
|
||||
|
||||
async def __call__(self, request_data: LiteLLM_AuditLogs) -> None:
|
||||
self.rows.append(request_data)
|
||||
|
||||
|
||||
def _kill_switch_app(
|
||||
role: LitellmUserRoles,
|
||||
http_client: _FakeKillSwitchClient,
|
||||
audit_log: _AuditLogRecorder | None = None,
|
||||
) -> TestClient:
|
||||
test_client: Final = _make_app_with_role(role)
|
||||
test_client.app.dependency_overrides[agent_endpoints.default_kill_switch_http_client] = lambda: http_client
|
||||
test_client.app.dependency_overrides[agent_endpoints.default_kill_switch_audit_log_writer] = (
|
||||
lambda: audit_log or _AuditLogRecorder()
|
||||
)
|
||||
return test_client
|
||||
|
||||
|
||||
def test_kill_switch_trigger_fires_the_configured_webhook_and_returns_the_result(monkeypatch) -> None:
|
||||
registry: Final = MagicMock()
|
||||
registry.get_agent_by_id = MagicMock(return_value=_agent_with_kill_switch())
|
||||
monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
|
||||
fake: Final = _FakeKillSwitchClient(httpx.Response(200, text="ok"))
|
||||
|
||||
resp: Final = _kill_switch_app(LitellmUserRoles.PROXY_ADMIN, fake).post(
|
||||
"/v1/agents/agent-123/kill_switch", headers={"Authorization": "Bearer k"}
|
||||
)
|
||||
|
||||
assert resp.status_code == 200, resp.text
|
||||
assert resp.json() == {
|
||||
"agent_id": "agent-123",
|
||||
"url": "https://ops.example.com/kill",
|
||||
"method": "POST",
|
||||
"status_code": 200,
|
||||
"response_body": "ok",
|
||||
"error": None,
|
||||
}
|
||||
(method, url, headers, body, _timeout) = fake.calls[0]
|
||||
assert (method, url, body) == ("POST", "https://ops.example.com/kill?reason=manual", {"action": "stop"})
|
||||
assert headers == {"X-Env": "prod", "Authorization": "Bearer tok-real"}
|
||||
|
||||
|
||||
def test_kill_switch_trigger_writes_an_audit_log_row_naming_the_admin_and_the_sanitized_result(monkeypatch) -> None:
|
||||
registry: Final = MagicMock()
|
||||
registry.get_agent_by_id = MagicMock(return_value=_agent_with_kill_switch())
|
||||
monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
|
||||
fake: Final = _FakeKillSwitchClient(httpx.Response(202, text='{"stopped": true}'))
|
||||
audit: Final = _AuditLogRecorder()
|
||||
test_client: Final = _kill_switch_app(LitellmUserRoles.PROXY_ADMIN, fake, audit)
|
||||
test_client.app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_id="test-user", user_role=LitellmUserRoles.PROXY_ADMIN, api_key="hashed-k"
|
||||
)
|
||||
|
||||
resp: Final = test_client.post("/v1/agents/agent-123/kill_switch", headers={"Authorization": "Bearer k"})
|
||||
|
||||
assert resp.status_code == 200, resp.text
|
||||
(row,) = audit.rows
|
||||
assert (row.action, row.table_name, row.object_id) == (
|
||||
"kill_switch_fired",
|
||||
LitellmTableNames.AGENT_TABLE_NAME,
|
||||
"agent-123",
|
||||
)
|
||||
assert (row.changed_by, row.changed_by_api_key) == ("test-user", "hashed-k")
|
||||
assert row.before_value is None
|
||||
assert json.loads(row.updated_values) == {
|
||||
"agent_id": "agent-123",
|
||||
"url": "https://ops.example.com/kill",
|
||||
"method": "POST",
|
||||
"status_code": 202,
|
||||
"response_body": '{"stopped": true}',
|
||||
}
|
||||
assert "tok-real" not in row.model_dump_json()
|
||||
|
||||
|
||||
def test_kill_switch_trigger_returns_502_and_still_audits_when_the_webhook_rejects(monkeypatch) -> None:
|
||||
registry: Final = MagicMock()
|
||||
registry.get_agent_by_id = MagicMock(return_value=_agent_with_kill_switch())
|
||||
monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
|
||||
fake: Final = _FakeKillSwitchClient(httpx.Response(401, text="bad token"))
|
||||
audit: Final = _AuditLogRecorder()
|
||||
|
||||
resp: Final = _kill_switch_app(LitellmUserRoles.PROXY_ADMIN, fake, audit).post(
|
||||
"/v1/agents/agent-123/kill_switch", headers={"Authorization": "Bearer k"}
|
||||
)
|
||||
|
||||
assert resp.status_code == 502, resp.text
|
||||
assert resp.json()["detail"]["status_code"] == 401
|
||||
assert resp.json()["detail"]["response_body"] == "bad token"
|
||||
(row,) = audit.rows
|
||||
assert row.action == "kill_switch_fired"
|
||||
assert json.loads(row.updated_values)["status_code"] == 401
|
||||
|
||||
|
||||
@pytest.mark.parametrize("role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY])
|
||||
def test_kill_switch_trigger_is_refused_before_any_webhook_call_for_non_admins(monkeypatch, role) -> None:
|
||||
registry: Final = MagicMock()
|
||||
registry.get_agent_by_id = MagicMock(return_value=_agent_with_kill_switch())
|
||||
monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
|
||||
fake: Final = _FakeKillSwitchClient(httpx.Response(200))
|
||||
audit: Final = _AuditLogRecorder()
|
||||
|
||||
resp: Final = _kill_switch_app(role, fake, audit).post(
|
||||
"/v1/agents/agent-123/kill_switch", headers={"Authorization": "Bearer k"}
|
||||
)
|
||||
|
||||
assert resp.status_code == 403, resp.text
|
||||
assert fake.calls == []
|
||||
assert audit.rows == []
|
||||
|
||||
|
||||
def test_kill_switch_trigger_404s_unknown_agent_and_400s_an_agent_without_one(monkeypatch) -> None:
|
||||
registry: Final = MagicMock()
|
||||
registry.get_agent_by_id = MagicMock(side_effect=[None, _sample_agent_response()])
|
||||
monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
|
||||
fake: Final = _FakeKillSwitchClient(httpx.Response(200))
|
||||
audit: Final = _AuditLogRecorder()
|
||||
test_client: Final = _kill_switch_app(LitellmUserRoles.PROXY_ADMIN, fake, audit)
|
||||
|
||||
missing: Final = test_client.post("/v1/agents/nope/kill_switch", headers={"Authorization": "Bearer k"})
|
||||
unconfigured: Final = test_client.post("/v1/agents/agent-123/kill_switch", headers={"Authorization": "Bearer k"})
|
||||
|
||||
assert missing.status_code == 404
|
||||
assert unconfigured.status_code == 400
|
||||
assert "no kill_switch configured" in unconfigured.json()["detail"]
|
||||
assert fake.calls == []
|
||||
assert audit.rows == []
|
||||
|
||||
|
||||
def test_kill_switch_trigger_fires_the_db_row_config_over_a_stale_in_memory_copy(monkeypatch) -> None:
|
||||
"""Another replica may have updated the agent; the row is the source of truth for what gets fired."""
|
||||
registry: Final = MagicMock()
|
||||
registry.get_agent_by_id = MagicMock(return_value=_agent_with_kill_switch())
|
||||
monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry)
|
||||
db_row: Final = SimpleNamespace(
|
||||
agent_id="agent-123",
|
||||
kill_switch={"url": "https://ops.example.com/kill-v2", "method": "DELETE", "auth": None},
|
||||
)
|
||||
prisma: Final = MagicMock()
|
||||
prisma.db.litellm_agentstable.find_unique = AsyncMock(return_value=db_row)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma)
|
||||
fake: Final = _FakeKillSwitchClient(httpx.Response(204))
|
||||
|
||||
resp: Final = _kill_switch_app(LitellmUserRoles.PROXY_ADMIN, fake).post(
|
||||
"/v1/agents/agent-123/kill_switch", headers={"Authorization": "Bearer k"}
|
||||
)
|
||||
|
||||
assert resp.status_code == 200, resp.text
|
||||
(method, url, headers, body, _timeout) = fake.calls[0]
|
||||
assert (method, url, headers, body) == ("DELETE", "https://ops.example.com/kill-v2", {}, None)
|
||||
assert prisma.db.litellm_agentstable.find_unique.await_args.kwargs == {"where": {"agent_id": "agent-123"}}
|
||||
registry.get_agent_by_id.assert_not_called()
|
||||
|
||||
|
||||
def test_get_agent_redacts_kill_switch_secret_for_admins_and_hides_it_from_others(monkeypatch) -> None:
|
||||
registry: Final = MagicMock()
|
||||
registry.get_agent_by_id = MagicMock(return_value=_agent_with_kill_switch())
|
||||
registry.ids_for_agent = MagicMock(return_value=("agent-123",))
|
||||
monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", registry)
|
||||
|
||||
def _get_as(role: LitellmUserRoles):
|
||||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||||
mock_prisma.db.litellm_agentstable.find_unique = AsyncMock(return_value=None)
|
||||
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
|
||||
return _make_app_with_role(role).get("/v1/agents/agent-123", headers={"Authorization": "Bearer k"})
|
||||
|
||||
admin: Final = _get_as(LitellmUserRoles.PROXY_ADMIN)
|
||||
assert admin.status_code == 200, admin.text
|
||||
assert admin.json()["kill_switch"] == {
|
||||
**_KILL_SWITCH,
|
||||
"auth": {"type": "bearer", "token": REDACTED_BY_LITELM_STRING},
|
||||
}
|
||||
|
||||
internal: Final = _get_as(LitellmUserRoles.INTERNAL_USER)
|
||||
assert internal.status_code == 200, internal.text
|
||||
assert internal.json()["kill_switch"] is None
|
||||
assert "tok-real" not in internal.text
|
||||
|
|
|
|||
248
tests/test_litellm/proxy/agent_endpoints/test_kill_switch.py
Normal file
248
tests/test_litellm/proxy/agent_endpoints/test_kill_switch.py
Normal file
|
|
@ -0,0 +1,248 @@
|
|||
from base64 import b64encode
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.constants import REDACTED_BY_LITELM_STRING
|
||||
from litellm.proxy.agent_endpoints.kill_switch import (
|
||||
build_kill_switch_request,
|
||||
fire_kill_switch,
|
||||
redact_kill_switch,
|
||||
restore_kill_switch,
|
||||
)
|
||||
from litellm.types.agents import AgentKillSwitchConfig
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _SentRequest:
|
||||
method: str
|
||||
url: str
|
||||
headers: Mapping[str, str]
|
||||
json: Mapping[str, object] | None
|
||||
timeout: float
|
||||
|
||||
|
||||
class _RecordingClient:
|
||||
def __init__(self, respond: httpx.Response | httpx.HTTPError) -> None:
|
||||
self.sent: list[_SentRequest] = [] # mutable-ok: test double records calls
|
||||
self.follow_redirects: list[bool] = [] # mutable-ok: test double records calls
|
||||
self._respond: Final = respond
|
||||
|
||||
def build_request(
|
||||
self,
|
||||
method: str,
|
||||
url: str,
|
||||
*,
|
||||
headers: Mapping[str, str],
|
||||
json: Mapping[str, object] | None,
|
||||
timeout: float,
|
||||
) -> httpx.Request:
|
||||
self.sent.append(_SentRequest(method, url, headers, json, timeout))
|
||||
return httpx.Request(method, url, headers=dict(headers), json=json)
|
||||
|
||||
async def send(self, request: httpx.Request, *, stream: bool, follow_redirects: bool) -> httpx.Response:
|
||||
self.follow_redirects.append(follow_redirects)
|
||||
if isinstance(self._respond, httpx.HTTPError):
|
||||
raise self._respond
|
||||
return self._respond
|
||||
|
||||
|
||||
class _CountingStream(httpx.AsyncByteStream):
|
||||
def __init__(self, chunk: bytes, chunks: int) -> None:
|
||||
self.pulled: int = 0 # rebind-ok: test double counts reads
|
||||
self._chunk: Final = chunk
|
||||
self._chunks: Final = chunks
|
||||
|
||||
async def __aiter__(self):
|
||||
for _ in range(self._chunks):
|
||||
self.pulled += 1 # rebind-ok: test double counts reads
|
||||
yield self._chunk
|
||||
|
||||
|
||||
def _config(**overrides: object) -> AgentKillSwitchConfig:
|
||||
return AgentKillSwitchConfig.model_validate({"url": "https://ops.example.com/agents/kill", **overrides})
|
||||
|
||||
|
||||
def test_request_carries_endpoint_method_query_params_headers_and_body() -> None:
|
||||
request: Final = build_kill_switch_request(
|
||||
_config(
|
||||
url="https://ops.example.com/kill?env=prod",
|
||||
method="PUT",
|
||||
query_params={"agent": "billing-bot", "reason": "manual stop"},
|
||||
headers={"X-Trace": "abc"},
|
||||
body={"action": "stop", "hard": True},
|
||||
)
|
||||
)
|
||||
|
||||
assert request.method == "PUT"
|
||||
assert str(httpx.URL(request.url)) == "https://ops.example.com/kill?env=prod&agent=billing-bot&reason=manual+stop"
|
||||
assert dict(request.headers) == {"X-Trace": "abc"}
|
||||
assert request.json_body == {"action": "stop", "hard": True}
|
||||
|
||||
|
||||
def test_request_defaults_to_post_with_no_body_and_untouched_url() -> None:
|
||||
request: Final = build_kill_switch_request(_config())
|
||||
|
||||
assert (request.method, request.url, dict(request.headers), request.json_body) == (
|
||||
"POST",
|
||||
"https://ops.example.com/agents/kill",
|
||||
{},
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("auth", "expected_headers"),
|
||||
[
|
||||
({"type": "bearer", "token": "tok-123"}, {"Authorization": "Bearer tok-123"}),
|
||||
({"type": "api_key", "api_key": "k-456"}, {"x-api-key": "k-456"}),
|
||||
({"type": "api_key", "header_name": "X-Ops-Key", "api_key": "k-456"}, {"X-Ops-Key": "k-456"}),
|
||||
(
|
||||
{"type": "basic", "username": "ops", "password": "pw:1"},
|
||||
{"Authorization": f"Basic {b64encode(b'ops:pw:1').decode()}"},
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_auth_becomes_the_matching_request_header(auth: Mapping[str, object], expected_headers: dict[str, str]) -> None:
|
||||
request: Final = build_kill_switch_request(_config(auth=auth))
|
||||
|
||||
assert dict(request.headers) == expected_headers
|
||||
|
||||
|
||||
def test_auth_header_wins_over_a_conflicting_custom_header() -> None:
|
||||
request: Final = build_kill_switch_request(
|
||||
_config(headers={"Authorization": "stale", "X-Env": "prod"}, auth={"type": "bearer", "token": "fresh"})
|
||||
)
|
||||
|
||||
assert dict(request.headers) == {"Authorization": "Bearer fresh", "X-Env": "prod"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("url", ["ftp://ops.example.com/kill", "/relative/kill", "ops.example.com/kill", ""])
|
||||
def test_config_rejects_non_http_urls(url: str) -> None:
|
||||
with pytest.raises(ValidationError, match="absolute http"):
|
||||
_config(url=url)
|
||||
|
||||
|
||||
def test_config_rejects_unknown_auth_type_and_unknown_fields() -> None:
|
||||
with pytest.raises(ValidationError):
|
||||
_config(auth={"type": "hmac", "secret": "x"})
|
||||
with pytest.raises(ValidationError):
|
||||
_config(endpoint="https://typo.example.com")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("auth", "secret_field"),
|
||||
[
|
||||
({"type": "bearer", "token": "tok-123"}, "token"),
|
||||
({"type": "api_key", "header_name": "X-K", "api_key": "k-456"}, "api_key"),
|
||||
({"type": "basic", "username": "ops", "password": "pw"}, "password"),
|
||||
],
|
||||
)
|
||||
def test_redact_replaces_only_the_secret_and_restore_puts_it_back(auth: dict[str, str], secret_field: str) -> None:
|
||||
original: Final = _config(auth=auth)
|
||||
|
||||
redacted: Final = redact_kill_switch(original)
|
||||
assert redacted is not None and redacted.auth is not None
|
||||
assert redacted.auth.model_dump() == {**auth, secret_field: REDACTED_BY_LITELM_STRING}
|
||||
assert original.auth is not None and original.auth.model_dump() == auth, "redact must not mutate its input"
|
||||
|
||||
restored: Final = restore_kill_switch(redacted, original)
|
||||
assert restored == original
|
||||
|
||||
|
||||
def test_restore_keeps_a_rotated_secret_and_never_stores_the_marker_itself() -> None:
|
||||
rotated: Final = _config(auth={"type": "bearer", "token": "new-token"})
|
||||
stored: Final = _config(auth={"type": "bearer", "token": "old-token"})
|
||||
assert restore_kill_switch(rotated, stored) == rotated
|
||||
assert restore_kill_switch(None, stored) is None
|
||||
|
||||
marker_only: Final = _config(auth={"type": "bearer", "token": REDACTED_BY_LITELM_STRING})
|
||||
assert restore_kill_switch(marker_only, None) == _config(auth={"type": "bearer", "token": ""})
|
||||
|
||||
|
||||
def test_restore_does_not_borrow_a_secret_from_a_different_auth_type() -> None:
|
||||
incoming: Final = _config(auth={"type": "bearer", "token": REDACTED_BY_LITELM_STRING})
|
||||
stored: Final = _config(auth={"type": "api_key", "api_key": "k-456"})
|
||||
|
||||
assert restore_kill_switch(incoming, stored) == _config(auth={"type": "bearer", "token": ""})
|
||||
|
||||
|
||||
def test_redact_passes_through_configs_without_auth() -> None:
|
||||
assert redact_kill_switch(None) is None
|
||||
plain: Final = _config(headers={"X-Env": "prod"})
|
||||
assert redact_kill_switch(plain) is plain
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fire_sends_exactly_the_built_request_and_reports_the_2xx_reply_without_the_query() -> None:
|
||||
client: Final = _RecordingClient(httpx.Response(202, text="stopping"))
|
||||
config: Final = _config(
|
||||
method="DELETE",
|
||||
query_params={"force": "1", "token": "qs-secret"},
|
||||
headers={"X-Env": "prod"},
|
||||
body={"agent": "billing-bot"},
|
||||
auth={"type": "bearer", "token": "tok-123"},
|
||||
)
|
||||
|
||||
result: Final = await fire_kill_switch(agent_id="agent-1", config=config, http_client=client, timeout=3.5)
|
||||
|
||||
assert client.sent == [
|
||||
_SentRequest(
|
||||
method="DELETE",
|
||||
url="https://ops.example.com/agents/kill?force=1&token=qs-secret",
|
||||
headers={"X-Env": "prod", "Authorization": "Bearer tok-123"},
|
||||
json={"agent": "billing-bot"},
|
||||
timeout=3.5,
|
||||
)
|
||||
]
|
||||
assert client.follow_redirects == [False], "a redirecting webhook must not be followed to another host"
|
||||
assert result.succeeded is True
|
||||
assert result.model_dump() == {
|
||||
"agent_id": "agent-1",
|
||||
"url": "https://ops.example.com/agents/kill",
|
||||
"method": "DELETE",
|
||||
"status_code": 202,
|
||||
"response_body": "stopping",
|
||||
"error": None,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fire_reports_a_non_2xx_reply_as_failure_with_the_body() -> None:
|
||||
client: Final = _RecordingClient(httpx.Response(503, text="x" * 5000))
|
||||
|
||||
result: Final = await fire_kill_switch(agent_id="agent-1", config=_config(), http_client=client)
|
||||
|
||||
assert result.succeeded is False
|
||||
assert result.status_code == 503
|
||||
assert result.response_body == "x" * 2000
|
||||
assert result.error is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fire_stops_reading_the_body_at_the_cap_instead_of_buffering_the_whole_reply() -> None:
|
||||
stream: Final = _CountingStream(b"y" * 500, chunks=100)
|
||||
client: Final = _RecordingClient(httpx.Response(200, stream=stream))
|
||||
|
||||
result: Final = await fire_kill_switch(agent_id="agent-1", config=_config(), http_client=client)
|
||||
|
||||
assert result.response_body == "y" * 2000
|
||||
assert stream.pulled == 4, f"read {stream.pulled} of 100 chunks for a 2000 char cap"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fire_reports_a_transport_error_by_type_without_raising_or_echoing_the_url() -> None:
|
||||
client: Final = _RecordingClient(httpx.ConnectError("boom https://ops.example.com/agents/kill?token=qs-secret"))
|
||||
|
||||
result: Final = await fire_kill_switch(
|
||||
agent_id="agent-1", config=_config(query_params={"token": "qs-secret"}), http_client=client
|
||||
)
|
||||
|
||||
assert result.succeeded is False
|
||||
assert (result.status_code, result.response_body) == (None, None)
|
||||
assert result.error == "ConnectError"
|
||||
assert "qs-secret" not in result.model_dump_json()
|
||||
|
|
@ -3776,6 +3776,7 @@ AGENT_MANAGEMENT_ROUTES = [
|
|||
"/v1/agents/abc-123",
|
||||
"/v1/agents/make_public",
|
||||
"/v1/agents/abc-123/make_public",
|
||||
"/v1/agents/abc-123/kill_switch",
|
||||
]
|
||||
|
||||
AGENT_INFERENCE_ROUTES = [
|
||||
|
|
|
|||
|
|
@ -0,0 +1,439 @@
|
|||
import json
|
||||
import re
|
||||
from collections.abc import Mapping
|
||||
from datetime import date, datetime, timezone
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import httpx
|
||||
import psycopg
|
||||
import pytest
|
||||
from psycopg.rows import dict_row
|
||||
from pydantic import ValidationError
|
||||
from pytest_postgresql import factories
|
||||
|
||||
from litellm.constants import (
|
||||
SPEND_CAPTURE_RATE_CHECK_JOB_ID,
|
||||
SPEND_CAPTURE_RATE_DOCS_URL,
|
||||
SPEND_CAPTURE_RATE_MAX_RANGE_DAYS,
|
||||
)
|
||||
from litellm.llms.openai.organization_costs import OPENAI_ADMIN_KEY_ENV_VAR
|
||||
from litellm.proxy.spend_tracking.spend_capture_rate import (
|
||||
ProviderBillingCredentialMissing,
|
||||
ProviderBillingRequestFailed,
|
||||
alert_message,
|
||||
captured_spend_by_day,
|
||||
compute_capture_rate,
|
||||
run_scheduled_spend_capture_rate_check,
|
||||
run_spend_capture_rate_check,
|
||||
)
|
||||
from litellm.types.proxy.spend_capture_rate import CaptureRateReport, SpendCaptureRateCheckSettings
|
||||
|
||||
_ADMIN_KEY: Final = "sk-admin-test"
|
||||
|
||||
|
||||
def _utc_midnight(day: str) -> int:
|
||||
return int(datetime.fromisoformat(day).replace(tzinfo=timezone.utc).timestamp())
|
||||
|
||||
|
||||
def _bucket(day: str, *amounts: float) -> dict[str, object]:
|
||||
return {
|
||||
"object": "bucket",
|
||||
"start_time": _utc_midnight(day),
|
||||
"end_time": _utc_midnight(day) + 86400,
|
||||
"results": [
|
||||
{"object": "organization.costs.result", "amount": {"value": amount, "currency": "usd"}, "line_item": None}
|
||||
for amount in amounts
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
class _FakeCostsApi:
|
||||
"""Serves ``pages`` in order and records every request it saw."""
|
||||
|
||||
def __init__(self, *pages: dict[str, object] | httpx.Response) -> None:
|
||||
self.responses = [
|
||||
page if isinstance(page, httpx.Response) else httpx.Response(200, json=page) for page in pages
|
||||
]
|
||||
self.calls: list[tuple[str, Mapping[str, object], Mapping[str, str]]] = []
|
||||
|
||||
async def __call__(self, url: str, params: Mapping[str, object], headers: Mapping[str, str]) -> httpx.Response:
|
||||
self.calls.append((url, dict(params), dict(headers)))
|
||||
return self.responses[len(self.calls) - 1]
|
||||
|
||||
|
||||
def _page(*buckets: dict[str, object], next_page: str | None = None) -> dict[str, object]:
|
||||
return {"object": "page", "data": list(buckets), "has_more": next_page is not None, "next_page": next_page}
|
||||
|
||||
|
||||
def _fake_prisma(rows: list[dict[str, object]]) -> MagicMock:
|
||||
prisma = MagicMock()
|
||||
prisma.db.query_raw = AsyncMock(return_value=rows)
|
||||
return prisma
|
||||
|
||||
|
||||
def test_capture_rate_covers_every_day_in_the_range_and_flags_the_threshold():
|
||||
report = compute_capture_rate(
|
||||
provider="openai",
|
||||
start_date=date(2026, 9, 20),
|
||||
end_date=date(2026, 9, 22),
|
||||
captured_by_day={"2026-09-20": 8.0, "2026-09-22": 1.0},
|
||||
billed_by_day={"2026-09-20": 10.0, "2026-09-21": 5.0},
|
||||
threshold=0.9,
|
||||
)
|
||||
|
||||
assert [day.date for day in report.days] == ["2026-09-20", "2026-09-21", "2026-09-22"]
|
||||
assert [day.capture_rate for day in report.days] == [0.8, 0.0, None]
|
||||
assert report.captured_spend == 9.0
|
||||
assert report.provider_spend == 15.0
|
||||
assert report.capture_rate == 0.6
|
||||
assert report.below_threshold is True
|
||||
|
||||
|
||||
def test_capture_rate_at_or_above_the_threshold_is_not_flagged_and_a_zero_bill_has_no_rate():
|
||||
healthy = compute_capture_rate(
|
||||
provider="openai",
|
||||
start_date=date(2026, 9, 20),
|
||||
end_date=date(2026, 9, 20),
|
||||
captured_by_day={"2026-09-20": 9.5},
|
||||
billed_by_day={"2026-09-20": 10.0},
|
||||
threshold=0.9,
|
||||
)
|
||||
over = compute_capture_rate(
|
||||
provider="openai",
|
||||
start_date=date(2026, 9, 20),
|
||||
end_date=date(2026, 9, 20),
|
||||
captured_by_day={"2026-09-20": 12.0},
|
||||
billed_by_day={"2026-09-20": 10.0},
|
||||
threshold=0.9,
|
||||
)
|
||||
unbilled = compute_capture_rate(
|
||||
provider="openai",
|
||||
start_date=date(2026, 9, 20),
|
||||
end_date=date(2026, 9, 20),
|
||||
captured_by_day={"2026-09-20": 3.0},
|
||||
billed_by_day={},
|
||||
threshold=0.9,
|
||||
)
|
||||
|
||||
assert (healthy.capture_rate, healthy.below_threshold) == (0.95, False)
|
||||
assert (over.capture_rate, over.below_threshold) == (1.2, False)
|
||||
assert (unbilled.capture_rate, unbilled.below_threshold) == (None, False)
|
||||
assert alert_message(healthy) is None
|
||||
assert alert_message(over) is None
|
||||
assert alert_message(unbilled) is None
|
||||
|
||||
|
||||
def test_alert_messages_name_the_cause_and_link_the_docs():
|
||||
below = compute_capture_rate(
|
||||
provider="openai",
|
||||
start_date=date(2026, 9, 14),
|
||||
end_date=date(2026, 9, 20),
|
||||
captured_by_day={"2026-09-14": 700.0},
|
||||
billed_by_day={"2026-09-14": 1000.0},
|
||||
threshold=0.9,
|
||||
)
|
||||
|
||||
below_message = alert_message(below)
|
||||
missing_message = alert_message(ProviderBillingCredentialMissing("openai", OPENAI_ADMIN_KEY_ENV_VAR))
|
||||
failed_message = alert_message(ProviderBillingRequestFailed("openai", "HTTP 401: nope"))
|
||||
|
||||
assert below_message is not None and "70.0%" in below_message and "90%" in below_message
|
||||
assert "$700.00" in below_message and "$1,000.00" in below_message
|
||||
assert "2026-09-14 to 2026-09-20" in below_message
|
||||
assert missing_message is not None and OPENAI_ADMIN_KEY_ENV_VAR in missing_message
|
||||
assert failed_message is not None and "HTTP 401: nope" in failed_message
|
||||
assert all(SPEND_CAPTURE_RATE_DOCS_URL in m for m in (below_message, missing_message, failed_message))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_reads_the_closed_window_before_today_and_publishes_the_rate(monkeypatch):
|
||||
monkeypatch.setenv(OPENAI_ADMIN_KEY_ENV_VAR, _ADMIN_KEY)
|
||||
api = _FakeCostsApi(_page(_bucket("2026-09-21", 100.0), _bucket("2026-09-22", 100.0)))
|
||||
prisma = _fake_prisma([{"date": "2026-09-21", "spend": 95.0}, {"date": "2026-09-22", "spend": 91.0}])
|
||||
alert = AsyncMock()
|
||||
publish = MagicMock()
|
||||
|
||||
results = await run_spend_capture_rate_check(
|
||||
prisma,
|
||||
SpendCaptureRateCheckSettings(lookback_days=2, threshold=0.9),
|
||||
alert=alert,
|
||||
publish=publish,
|
||||
today=date(2026, 9, 23),
|
||||
http_get=api,
|
||||
)
|
||||
|
||||
(report,) = results
|
||||
assert isinstance(report, CaptureRateReport)
|
||||
assert (report.start_date, report.end_date) == ("2026-09-21", "2026-09-22")
|
||||
assert report.capture_rate == 0.93
|
||||
publish.assert_called_once_with("openai", 0.93)
|
||||
alert.assert_not_awaited()
|
||||
assert api.calls[0][1]["start_time"] == _utc_midnight("2026-09-21")
|
||||
assert api.calls[0][1]["end_time"] == _utc_midnight("2026-09-23")
|
||||
sql, start, end, providers = prisma.db.query_raw.await_args.args
|
||||
assert (start, end) == ("2026-09-21", "2026-09-22")
|
||||
assert set(providers) == {"openai", "text-completion-openai"}
|
||||
assert '"LiteLLM_DailyUserSpend"' in sql
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_alerts_under_the_threshold_and_still_publishes_the_rate(monkeypatch):
|
||||
monkeypatch.setenv(OPENAI_ADMIN_KEY_ENV_VAR, _ADMIN_KEY)
|
||||
api = _FakeCostsApi(_page(_bucket("2026-09-22", 200.0)))
|
||||
prisma = _fake_prisma([{"date": "2026-09-22", "spend": 50.0}])
|
||||
alert = AsyncMock()
|
||||
publish = MagicMock()
|
||||
|
||||
await run_spend_capture_rate_check(
|
||||
prisma,
|
||||
SpendCaptureRateCheckSettings(lookback_days=1),
|
||||
alert=alert,
|
||||
publish=publish,
|
||||
today=date(2026, 9, 23),
|
||||
http_get=api,
|
||||
)
|
||||
|
||||
publish.assert_called_once_with("openai", 0.25)
|
||||
alert.assert_awaited_once()
|
||||
assert "25.0%" in alert.await_args.args[0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_alerts_on_a_missing_admin_key_and_publishes_nothing(monkeypatch):
|
||||
monkeypatch.delenv(OPENAI_ADMIN_KEY_ENV_VAR, raising=False)
|
||||
api = _FakeCostsApi()
|
||||
prisma = _fake_prisma([])
|
||||
alert = AsyncMock()
|
||||
publish = MagicMock()
|
||||
|
||||
(result,) = await run_spend_capture_rate_check(
|
||||
prisma, SpendCaptureRateCheckSettings(), alert=alert, publish=publish, today=date(2026, 9, 23), http_get=api
|
||||
)
|
||||
|
||||
assert result == ProviderBillingCredentialMissing("openai", OPENAI_ADMIN_KEY_ENV_VAR)
|
||||
publish.assert_called_once_with("openai", None)
|
||||
alert.assert_awaited_once()
|
||||
assert api.calls == []
|
||||
prisma.db.query_raw.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_alerts_on_an_unreadable_bill_and_publishes_no_rate(monkeypatch):
|
||||
monkeypatch.setenv(OPENAI_ADMIN_KEY_ENV_VAR, _ADMIN_KEY)
|
||||
api = _FakeCostsApi(httpx.Response(401, json={"error": {"message": "Incorrect API key provided"}}))
|
||||
prisma = _fake_prisma([{"date": "2026-09-22", "spend": 5.0}])
|
||||
alert = AsyncMock()
|
||||
publish = MagicMock()
|
||||
|
||||
(result,) = await run_spend_capture_rate_check(
|
||||
prisma, SpendCaptureRateCheckSettings(), alert=alert, publish=publish, today=date(2026, 9, 23), http_get=api
|
||||
)
|
||||
|
||||
assert result == ProviderBillingRequestFailed("openai", "HTTP 401: " + api.responses[0].text[:300])
|
||||
publish.assert_called_once_with("openai", None)
|
||||
alert.assert_awaited_once()
|
||||
assert "HTTP 401" in alert.await_args.args[0]
|
||||
prisma.db.query_raw.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_publishes_no_rate_when_the_provider_billed_nothing(monkeypatch):
|
||||
monkeypatch.setenv(OPENAI_ADMIN_KEY_ENV_VAR, _ADMIN_KEY)
|
||||
api = _FakeCostsApi(_page())
|
||||
prisma = _fake_prisma([{"date": "2026-09-22", "spend": 5.0}])
|
||||
publish = MagicMock()
|
||||
alert = AsyncMock()
|
||||
|
||||
(report,) = await run_spend_capture_rate_check(
|
||||
prisma,
|
||||
SpendCaptureRateCheckSettings(lookback_days=1),
|
||||
alert=alert,
|
||||
publish=publish,
|
||||
today=date(2026, 9, 23),
|
||||
http_get=api,
|
||||
)
|
||||
|
||||
assert isinstance(report, CaptureRateReport) and report.capture_rate is None
|
||||
publish.assert_called_once_with("openai", None)
|
||||
alert.assert_not_awaited()
|
||||
|
||||
|
||||
def _pod_lock(acquired: bool) -> MagicMock:
|
||||
lock = MagicMock()
|
||||
lock.redis_cache = MagicMock()
|
||||
lock.redis_cache.async_get_cache = AsyncMock(return_value="other-pod")
|
||||
lock.get_redis_lock_key = MagicMock(return_value="lock-key")
|
||||
lock.acquire_lock = AsyncMock(return_value=acquired)
|
||||
lock.release_lock = AsyncMock()
|
||||
return lock
|
||||
|
||||
|
||||
async def _scheduled_run(
|
||||
lock: MagicMock, monkeypatch, *, captured: float, prisma: MagicMock | None = None
|
||||
) -> tuple[AsyncMock, MagicMock]:
|
||||
monkeypatch.setenv(OPENAI_ADMIN_KEY_ENV_VAR, _ADMIN_KEY)
|
||||
alert = AsyncMock()
|
||||
publish = MagicMock()
|
||||
await run_scheduled_spend_capture_rate_check(
|
||||
prisma or _fake_prisma([{"date": "2026-09-22", "spend": captured}]),
|
||||
SpendCaptureRateCheckSettings(lookback_days=1),
|
||||
pod_lock_manager=lock,
|
||||
alert=alert,
|
||||
publish=publish,
|
||||
today=date(2026, 9, 23),
|
||||
http_get=_FakeCostsApi(_page(_bucket("2026-09-22", 200.0))),
|
||||
)
|
||||
return alert, publish
|
||||
|
||||
|
||||
async def _scheduled_run_under_threshold(lock: MagicMock, monkeypatch) -> tuple[AsyncMock, MagicMock]:
|
||||
return await _scheduled_run(lock, monkeypatch, captured=50.0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_healthy_scheduled_check_publishes_and_never_touches_the_alert_lock(monkeypatch):
|
||||
lock = _pod_lock(acquired=False)
|
||||
|
||||
alert, publish = await _scheduled_run(lock, monkeypatch, captured=190.0)
|
||||
|
||||
publish.assert_called_once_with("openai", 0.95)
|
||||
alert.assert_not_awaited()
|
||||
lock.acquire_lock.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_scheduled_check_that_fails_never_claims_the_alert_window(monkeypatch):
|
||||
lock = _pod_lock(acquired=True)
|
||||
prisma = MagicMock()
|
||||
prisma.db.query_raw = AsyncMock(side_effect=RuntimeError("database gone"))
|
||||
|
||||
with pytest.raises(RuntimeError, match="database gone"):
|
||||
await _scheduled_run(lock, monkeypatch, captured=0.0, prisma=prisma)
|
||||
|
||||
lock.acquire_lock.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scheduled_check_publishes_but_stays_quiet_when_another_pod_holds_the_alert_window(monkeypatch):
|
||||
lock = _pod_lock(acquired=False)
|
||||
|
||||
alert, publish = await _scheduled_run_under_threshold(lock, monkeypatch)
|
||||
|
||||
publish.assert_called_once_with("openai", 0.25)
|
||||
alert.assert_not_awaited()
|
||||
lock.release_lock.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scheduled_check_alerts_and_keeps_the_lock_until_it_expires_when_it_wins(monkeypatch):
|
||||
lock = _pod_lock(acquired=True)
|
||||
|
||||
alert, publish = await _scheduled_run_under_threshold(lock, monkeypatch)
|
||||
|
||||
lock.acquire_lock.assert_awaited_once_with(cronjob_id=SPEND_CAPTURE_RATE_CHECK_JOB_ID, ttl=900)
|
||||
lock.release_lock.assert_not_awaited()
|
||||
publish.assert_called_once_with("openai", 0.25)
|
||||
alert.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scheduled_check_alerts_when_the_lock_cannot_be_acquired_or_read(monkeypatch):
|
||||
lock = _pod_lock(acquired=False)
|
||||
lock.redis_cache.async_get_cache = AsyncMock(side_effect=ConnectionError("redis down"))
|
||||
|
||||
alert, publish = await _scheduled_run_under_threshold(lock, monkeypatch)
|
||||
|
||||
publish.assert_called_once_with("openai", 0.25)
|
||||
alert.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_scheduled_check_alerts_without_a_lock_manager(monkeypatch):
|
||||
lock = _pod_lock(acquired=False)
|
||||
lock.redis_cache = None
|
||||
|
||||
alert, publish = await _scheduled_run_under_threshold(lock, monkeypatch)
|
||||
|
||||
lock.acquire_lock.assert_not_awaited()
|
||||
publish.assert_called_once_with("openai", 0.25)
|
||||
alert.assert_awaited_once()
|
||||
|
||||
|
||||
def test_settings_reject_typos_and_out_of_range_values():
|
||||
with pytest.raises(ValidationError, match="threshhold"):
|
||||
SpendCaptureRateCheckSettings.model_validate({"threshhold": 0.9})
|
||||
with pytest.raises(ValidationError, match="threshold"):
|
||||
SpendCaptureRateCheckSettings.model_validate({"threshold": 1.5})
|
||||
with pytest.raises(ValidationError, match="providers"):
|
||||
SpendCaptureRateCheckSettings.model_validate({"providers": []})
|
||||
with pytest.raises(ValidationError, match="providers"):
|
||||
SpendCaptureRateCheckSettings.model_validate({"providers": ["anthropic"]})
|
||||
with pytest.raises(ValidationError, match="lookback_days"):
|
||||
SpendCaptureRateCheckSettings.model_validate({"lookback_days": SPEND_CAPTURE_RATE_MAX_RANGE_DAYS + 1})
|
||||
parsed = SpendCaptureRateCheckSettings.model_validate(
|
||||
json.loads('{"providers": ["openai"], "threshold": 0.8, "lookback_days": 3, "openai_project_ids": ["p"]}')
|
||||
)
|
||||
assert (parsed.threshold, parsed.lookback_days, parsed.openai_project_ids) == (0.8, 3, ("p",))
|
||||
|
||||
|
||||
_capture_postgresql_proc: Final = factories.postgresql_proc()
|
||||
_capture_postgresql: Final = factories.postgresql("_capture_postgresql_proc")
|
||||
|
||||
_DAILY_USER_SPEND_DDL: Final = """
|
||||
CREATE TABLE "LiteLLM_DailyUserSpend" (
|
||||
id TEXT PRIMARY KEY,
|
||||
date TEXT NOT NULL,
|
||||
custom_llm_provider TEXT,
|
||||
spend DOUBLE PRECISION DEFAULT 0
|
||||
)
|
||||
"""
|
||||
|
||||
|
||||
class _PsycopgPrisma:
|
||||
"""``prisma_client.db.query_raw`` on a real connection, with ``$n`` placeholders converted for psycopg."""
|
||||
|
||||
def __init__(self, conn: psycopg.Connection) -> None:
|
||||
self.db = self
|
||||
self._conn = conn
|
||||
|
||||
async def query_raw(self, sql: str, *params: object) -> list[dict[str, object]]:
|
||||
converted: Final = re.sub(r"\$(\d+)", r"%(p\1)s", sql)
|
||||
with self._conn.cursor(row_factory=dict_row) as cur:
|
||||
cur.execute(
|
||||
converted, # pyright: ignore[reportArgumentType] # psycopg stubs want a literal-typed query
|
||||
{f"p{i}": list(v) if isinstance(v, tuple) else v for i, v in enumerate(params, start=1)},
|
||||
)
|
||||
return cur.fetchall()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_captured_spend_sums_only_the_openai_billed_providers_inside_the_window(
|
||||
_capture_postgresql: psycopg.Connection,
|
||||
):
|
||||
conn: Final = _capture_postgresql
|
||||
conn.execute(_DAILY_USER_SPEND_DDL) # pyright: ignore[reportArgumentType] # DDL literal
|
||||
rows: Final = (
|
||||
("2026-09-19", "openai", 1.0),
|
||||
("2026-09-20", "openai", 2.0),
|
||||
("2026-09-20", "openai", 3.0),
|
||||
("2026-09-20", "text-completion-openai", 0.5),
|
||||
("2026-09-20", "anthropic", 100.0),
|
||||
("2026-09-21", "azure", 100.0),
|
||||
("2026-09-22", "openai", 4.0),
|
||||
)
|
||||
for index, (day, provider, spend) in enumerate(rows):
|
||||
conn.execute(
|
||||
'INSERT INTO "LiteLLM_DailyUserSpend" (id, date, custom_llm_provider, spend) VALUES (%s, %s, %s, %s)',
|
||||
(f"row-{index}", day, provider, spend),
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
captured = await captured_spend_by_day(
|
||||
_PsycopgPrisma(conn), # pyright: ignore[reportArgumentType] # duck-typed prisma for the raw query
|
||||
litellm_providers=("openai", "text-completion-openai"),
|
||||
start_date=date(2026, 9, 20),
|
||||
end_date=date(2026, 9, 21),
|
||||
)
|
||||
|
||||
assert dict(captured) == {"2026-09-20": 5.5}
|
||||
File diff suppressed because it is too large
Load diff
|
|
@ -14,7 +14,7 @@ from datetime import datetime, timedelta, timezone
|
|||
from pathlib import Path
|
||||
from typing import Final
|
||||
from unittest import mock
|
||||
from unittest.mock import AsyncMock, MagicMock, create_autospec, mock_open, patch
|
||||
from unittest.mock import AsyncMock, MagicMock, call, create_autospec, mock_open, patch
|
||||
|
||||
import click
|
||||
import fastapi.routing
|
||||
|
|
@ -3554,9 +3554,7 @@ async def test_load_config_without_role_permissions_leaves_every_role_unrestrict
|
|||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
|
||||
config_file: Final = tmp_path / "config.yaml"
|
||||
config_file.write_text(
|
||||
yaml.dump({"model_list": [], "general_settings": {"max_parallel_requests": 7}})
|
||||
)
|
||||
config_file.write_text(yaml.dump({"model_list": [], "general_settings": {"max_parallel_requests": 7}}))
|
||||
|
||||
_, _, settings = await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(config_file))
|
||||
|
||||
|
|
@ -3586,7 +3584,9 @@ async def test_load_config_rejects_malformed_role_permissions(tmp_path):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_load_config_compiles_key_alias_pattern_at_startup(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
async def test_load_config_compiles_key_alias_pattern_at_startup(
|
||||
tmp_path: Path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
|
||||
monkeypatch.setattr(litellm, "key_alias_pattern", None)
|
||||
|
|
@ -3622,9 +3622,7 @@ def test_os_environ_resolution_leaves_the_config_layer_holding_the_reference(mon
|
|||
assert resolved["general_settings"]["coordination_redis"]["password"] == "sk-nested-value"
|
||||
assert resolved["general_settings"]["master_key"] == "sk-nested-value"
|
||||
assert proxy_config.settings.config_value("master_key") == "os.environ/PROOF_NESTED_SECRET"
|
||||
assert proxy_config.settings.config_value("coordination_redis") == {
|
||||
"password": "os.environ/PROOF_NESTED_SECRET"
|
||||
}
|
||||
assert proxy_config.settings.config_value("coordination_redis") == {"password": "os.environ/PROOF_NESTED_SECRET"}
|
||||
|
||||
|
||||
def test_os_environ_resolution_reaches_dicts_nested_in_a_list(monkeypatch):
|
||||
|
|
@ -5814,7 +5812,9 @@ async def test_boot_warns_that_a_shadowed_database_value_will_never_apply(tmp_pa
|
|||
config_path.write_text(
|
||||
yaml.safe_dump({"model_list": [], "general_settings": {"allowed_ips": ["1.2.3.4"], "max_file_size_mb": 5}})
|
||||
)
|
||||
db_row: Final = types.SimpleNamespace(param_value={"allowed_ips": ["1.2.3.4", "5.6.7.8"], "max_parallel_requests": 7})
|
||||
db_row: Final = types.SimpleNamespace(
|
||||
param_value={"allowed_ips": ["1.2.3.4", "5.6.7.8"], "max_parallel_requests": 7}
|
||||
)
|
||||
|
||||
async def read_config_row(_prisma_client, param_name):
|
||||
return db_row if param_name == "general_settings" else None
|
||||
|
|
@ -8026,7 +8026,9 @@ async def test_update_general_settings_clearing_a_db_override_falls_back_to_the_
|
|||
proxy_config.settings.load_yaml({"maximum_spend_logs_cleanup_run_budget": "90s"})
|
||||
|
||||
with patch("litellm.proxy.proxy_server.general_settings", proxy_config.settings):
|
||||
await proxy_config._update_general_settings(db_general_settings={"maximum_spend_logs_cleanup_run_budget": "30s"})
|
||||
await proxy_config._update_general_settings(
|
||||
db_general_settings={"maximum_spend_logs_cleanup_run_budget": "30s"}
|
||||
)
|
||||
await proxy_config._update_general_settings(db_general_settings={"store_model_in_db": True})
|
||||
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
|
@ -8069,10 +8071,18 @@ async def test_update_general_settings_keeps_yaml_pass_through_endpoints_next_to
|
|||
request.query_params = {}
|
||||
return request
|
||||
|
||||
settings: Final = patch("litellm.proxy.proxy_server.general_settings", {"pass_through_endpoints": [yaml_endpoint]}) # test-quality-ok: the method reads this module global; no injection seam
|
||||
yaml_endpoints: Final = patch("litellm.proxy.proxy_server.config_passthrough_endpoints", [yaml_endpoint]) # test-quality-ok: module global holding the YAML endpoints the fix merges in
|
||||
initialize: Final = patch("litellm.proxy.proxy_server.initialize_pass_through_endpoints", AsyncMock()) # test-quality-ok: route registration needs the FastAPI app; auth is the observable here
|
||||
master_key: Final = patch("litellm.proxy.proxy_server.master_key", "sk-master") # test-quality-ok: a set master key is what makes a missing Authorization header a 401
|
||||
settings: Final = patch(
|
||||
"litellm.proxy.proxy_server.general_settings", {"pass_through_endpoints": [yaml_endpoint]}
|
||||
) # test-quality-ok: the method reads this module global; no injection seam
|
||||
yaml_endpoints: Final = patch(
|
||||
"litellm.proxy.proxy_server.config_passthrough_endpoints", [yaml_endpoint]
|
||||
) # test-quality-ok: module global holding the YAML endpoints the fix merges in
|
||||
initialize: Final = patch(
|
||||
"litellm.proxy.proxy_server.initialize_pass_through_endpoints", AsyncMock()
|
||||
) # test-quality-ok: route registration needs the FastAPI app; auth is the observable here
|
||||
master_key: Final = patch(
|
||||
"litellm.proxy.proxy_server.master_key", "sk-master"
|
||||
) # test-quality-ok: a set master key is what makes a missing Authorization header a 401
|
||||
with settings, yaml_endpoints, initialize, master_key:
|
||||
await ProxyConfig()._update_general_settings(db_general_settings={"pass_through_endpoints": [db_endpoint]})
|
||||
|
||||
|
|
@ -8119,10 +8129,18 @@ async def test_update_general_settings_db_pass_through_endpoint_cannot_override_
|
|||
request.headers = {}
|
||||
request.query_params = {}
|
||||
|
||||
settings: Final = patch("litellm.proxy.proxy_server.general_settings", {"pass_through_endpoints": [yaml_endpoint]}) # test-quality-ok: the method reads this module global; no injection seam
|
||||
yaml_endpoints: Final = patch("litellm.proxy.proxy_server.config_passthrough_endpoints", [yaml_endpoint]) # test-quality-ok: module global holding the YAML endpoints the fix merges in
|
||||
initialize: Final = patch("litellm.proxy.proxy_server.initialize_pass_through_endpoints", AsyncMock()) # test-quality-ok: route registration needs the FastAPI app; auth is the observable here
|
||||
master_key: Final = patch("litellm.proxy.proxy_server.master_key", "sk-master") # test-quality-ok: a set master key is what makes a missing Authorization header a 401
|
||||
settings: Final = patch(
|
||||
"litellm.proxy.proxy_server.general_settings", {"pass_through_endpoints": [yaml_endpoint]}
|
||||
) # test-quality-ok: the method reads this module global; no injection seam
|
||||
yaml_endpoints: Final = patch(
|
||||
"litellm.proxy.proxy_server.config_passthrough_endpoints", [yaml_endpoint]
|
||||
) # test-quality-ok: module global holding the YAML endpoints the fix merges in
|
||||
initialize: Final = patch(
|
||||
"litellm.proxy.proxy_server.initialize_pass_through_endpoints", AsyncMock()
|
||||
) # test-quality-ok: route registration needs the FastAPI app; auth is the observable here
|
||||
master_key: Final = patch(
|
||||
"litellm.proxy.proxy_server.master_key", "sk-master"
|
||||
) # test-quality-ok: a set master key is what makes a missing Authorization header a 401
|
||||
with settings, yaml_endpoints, initialize, master_key:
|
||||
await ProxyConfig()._update_general_settings(db_general_settings={"pass_through_endpoints": [db_endpoint]})
|
||||
|
||||
|
|
@ -8156,11 +8174,19 @@ async def test_deleting_the_stored_pass_through_row_takes_the_route_out_of_servi
|
|||
prior_registry: Final = dict(_registered_pass_through_routes)
|
||||
|
||||
def live_routes() -> set[str]:
|
||||
return {route for route in InitPassThroughEndpointHelpers.get_all_registered_pass_through_routes() if path in route}
|
||||
return {
|
||||
route for route in InitPassThroughEndpointHelpers.get_all_registered_pass_through_routes() if path in route
|
||||
}
|
||||
|
||||
settings: Final = patch("litellm.proxy.proxy_server.general_settings", {}) # test-quality-ok: the method reads this module global; no injection seam
|
||||
yaml_endpoints: Final = patch("litellm.proxy.proxy_server.config_passthrough_endpoints", None) # test-quality-ok: module global holding the YAML endpoints; this case has none
|
||||
app_routes: Final = patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.SafeRouteAdder.add_api_route_if_not_exists") # test-quality-ok: the registry is the observable; a real route would stay on the shared FastAPI app for the rest of the xdist worker
|
||||
settings: Final = patch(
|
||||
"litellm.proxy.proxy_server.general_settings", {}
|
||||
) # test-quality-ok: the method reads this module global; no injection seam
|
||||
yaml_endpoints: Final = patch(
|
||||
"litellm.proxy.proxy_server.config_passthrough_endpoints", None
|
||||
) # test-quality-ok: module global holding the YAML endpoints; this case has none
|
||||
app_routes: Final = patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.SafeRouteAdder.add_api_route_if_not_exists"
|
||||
) # test-quality-ok: the registry is the observable; a real route would stay on the shared FastAPI app for the rest of the xdist worker
|
||||
try:
|
||||
with settings, yaml_endpoints, app_routes:
|
||||
pc = ProxyConfig()
|
||||
|
|
@ -8201,9 +8227,15 @@ async def test_a_stored_pass_through_row_never_disturbs_the_config_declared_rout
|
|||
registered: Final = InitPassThroughEndpointHelpers.get_all_registered_pass_through_routes()
|
||||
return {path for path in (config_path, db_path) if any(path in route for route in registered)}
|
||||
|
||||
settings: Final = patch("litellm.proxy.proxy_server.general_settings", {"pass_through_endpoints": [config_endpoint]}) # test-quality-ok: the method reads this module global; no injection seam
|
||||
yaml_endpoints: Final = patch("litellm.proxy.proxy_server.config_passthrough_endpoints", [config_endpoint]) # test-quality-ok: module global holding the YAML endpoints the reload merges in
|
||||
app_routes: Final = patch("litellm.proxy.pass_through_endpoints.pass_through_endpoints.SafeRouteAdder.add_api_route_if_not_exists") # test-quality-ok: the registry is the observable; a real route would stay on the shared FastAPI app for the rest of the xdist worker
|
||||
settings: Final = patch(
|
||||
"litellm.proxy.proxy_server.general_settings", {"pass_through_endpoints": [config_endpoint]}
|
||||
) # test-quality-ok: the method reads this module global; no injection seam
|
||||
yaml_endpoints: Final = patch(
|
||||
"litellm.proxy.proxy_server.config_passthrough_endpoints", [config_endpoint]
|
||||
) # test-quality-ok: module global holding the YAML endpoints the reload merges in
|
||||
app_routes: Final = patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.SafeRouteAdder.add_api_route_if_not_exists"
|
||||
) # test-quality-ok: the registry is the observable; a real route would stay on the shared FastAPI app for the rest of the xdist worker
|
||||
try:
|
||||
with settings, yaml_endpoints, app_routes:
|
||||
await initialize_pass_through_endpoints(pass_through_endpoints=[config_endpoint])
|
||||
|
|
@ -9344,9 +9376,7 @@ async def test_increment_spend_counters_finalizes_after_unreserved_increments():
|
|||
async def assert_reservation_not_finalized_yet(**kwargs):
|
||||
assert budget_reservation["finalized"] is False
|
||||
incremented_counters.append(kwargs["counter_key"])
|
||||
return ps.PendingSpendIncrement(
|
||||
counter_key=kwargs["counter_key"], increment=kwargs["increment"]
|
||||
)
|
||||
return ps.PendingSpendIncrement(counter_key=kwargs["counter_key"], increment=kwargs["increment"])
|
||||
|
||||
import litellm.proxy.proxy_server as ps
|
||||
|
||||
|
|
@ -10943,9 +10973,15 @@ async def _lit6973_drive_realtime_session(
|
|||
side_effect=pre_call_error, return_value=({"model": "vertex_ai/gemini-live-2.5-flash"}, logging_obj)
|
||||
)
|
||||
ws: Final = websocket if websocket is not None else _lit6973_fake_realtime_ws()
|
||||
can_call = patch.object(ps, "can_key_call_resolved_model", new=AsyncMock(side_effect=model_access_error)) # test-quality-ok: no HTTP boundary; fakes in-process auth to reach the exit under test
|
||||
pre = patch.object(ps.ProxyBaseLLMRequestProcessing, "common_processing_pre_call_logic", new=pre_call) # test-quality-ok: fakes phase-1 wiring; assertion checks observable reservation state
|
||||
route = patch.object(ps, "route_request", new=AsyncMock(return_value=fake_llm_call())) # test-quality-ok: fakes the relay whose success/refusal outcome the endpoint reads off the logging object
|
||||
can_call = patch.object(
|
||||
ps, "can_key_call_resolved_model", new=AsyncMock(side_effect=model_access_error)
|
||||
) # test-quality-ok: no HTTP boundary; fakes in-process auth to reach the exit under test
|
||||
pre = patch.object(
|
||||
ps.ProxyBaseLLMRequestProcessing, "common_processing_pre_call_logic", new=pre_call
|
||||
) # test-quality-ok: fakes phase-1 wiring; assertion checks observable reservation state
|
||||
route = patch.object(
|
||||
ps, "route_request", new=AsyncMock(return_value=fake_llm_call())
|
||||
) # test-quality-ok: fakes the relay whose success/refusal outcome the endpoint reads off the logging object
|
||||
with can_call, pre, route:
|
||||
await ps.realtime_websocket_endpoint(
|
||||
websocket=ws,
|
||||
|
|
@ -11077,13 +11113,9 @@ async def _lit6463_drive_realtime_session_holding_a_max_parallel_slot(
|
|||
from litellm.proxy.utils import InternalUsageCache
|
||||
|
||||
dual_cache: Final = DualCache()
|
||||
await dual_cache.async_set_cache(
|
||||
key=_LIT6463_COUNTER_KEY, value={"slot-1": 1.0, "slot-2": 2.0}, local_only=True
|
||||
)
|
||||
await dual_cache.async_set_cache(key=_LIT6463_COUNTER_KEY, value={"slot-1": 1.0, "slot-2": 2.0}, local_only=True)
|
||||
limiter: Final = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=InternalUsageCache(dual_cache))
|
||||
stash: Final = RequestRateLimiterStash(
|
||||
parallel_slot={"slot_id": "slot-1", "counter_keys": [_LIT6463_COUNTER_KEY]}
|
||||
)
|
||||
stash: Final = RequestRateLimiterStash(parallel_slot={"slot_id": "slot-1", "counter_keys": [_LIT6463_COUNTER_KEY]})
|
||||
reservation: Final = {"reserved_cost": 0.55, "input_cost": 0.0, "finalized": False, "entries": []}
|
||||
|
||||
stash_token: Final = _request_stash.set(stash)
|
||||
|
|
@ -11135,9 +11167,7 @@ async def test_successful_realtime_session_leaves_the_max_parallel_slot_for_the_
|
|||
limiter's integer in-memory fallback, double-decrement the counter so the key
|
||||
admits more sessions than max_parallel_requests allows. With the success stamp
|
||||
present the route leaves the slot and the stash alone."""
|
||||
dual_cache, stash = await _lit6463_drive_realtime_session_holding_a_max_parallel_slot(
|
||||
backend_logged_success=True
|
||||
)
|
||||
dual_cache, stash = await _lit6463_drive_realtime_session_holding_a_max_parallel_slot(backend_logged_success=True)
|
||||
|
||||
assert await dual_cache.async_get_cache(key=_LIT6463_COUNTER_KEY, local_only=True) == {
|
||||
"slot-1": 1.0,
|
||||
|
|
@ -11183,8 +11213,12 @@ async def test_release_or_invalidate_falls_back_to_invalidating_the_counters():
|
|||
async def _record(counter_key: str) -> None:
|
||||
invalidated.append(counter_key)
|
||||
|
||||
failing_release = patch.object(br, "release_budget_reservation", new=AsyncMock(side_effect=RuntimeError("counter store down"))) # test-quality-ok: forces the failure branch; assertion observes which counter key got invalidated
|
||||
sink = patch.object(ps, "_invalidate_spend_counter", new=_record) # test-quality-ok: fakes the counter-store sink so the invalidated key is observable
|
||||
failing_release = patch.object(
|
||||
br, "release_budget_reservation", new=AsyncMock(side_effect=RuntimeError("counter store down"))
|
||||
) # test-quality-ok: forces the failure branch; assertion observes which counter key got invalidated
|
||||
sink = patch.object(
|
||||
ps, "_invalidate_spend_counter", new=_record
|
||||
) # test-quality-ok: fakes the counter-store sink so the invalidated key is observable
|
||||
with failing_release, sink:
|
||||
await br.release_or_invalidate_budget_reservation(budget_reservation=reservation)
|
||||
|
||||
|
|
@ -11200,8 +11234,12 @@ async def test_release_or_invalidate_finalizes_even_when_the_invalidate_fallback
|
|||
from litellm.proxy.spend_tracking import budget_reservation as br
|
||||
|
||||
reservation: Final = {"reserved_cost": 0.55, "input_cost": 0.0, "finalized": False, "entries": []}
|
||||
failing_release = patch.object(br, "release_budget_reservation", new=AsyncMock(side_effect=RuntimeError("counter store down"))) # test-quality-ok: forces the fallback branch
|
||||
failing_invalidate = patch.object(br, "invalidate_budget_reservation_counters", new=AsyncMock(side_effect=RuntimeError("still down"))) # test-quality-ok: forces the fallback itself to fail
|
||||
failing_release = patch.object(
|
||||
br, "release_budget_reservation", new=AsyncMock(side_effect=RuntimeError("counter store down"))
|
||||
) # test-quality-ok: forces the fallback branch
|
||||
failing_invalidate = patch.object(
|
||||
br, "invalidate_budget_reservation_counters", new=AsyncMock(side_effect=RuntimeError("still down"))
|
||||
) # test-quality-ok: forces the fallback itself to fail
|
||||
|
||||
with failing_release, failing_invalidate:
|
||||
await br.release_or_invalidate_budget_reservation(budget_reservation=reservation)
|
||||
|
|
@ -11257,6 +11295,49 @@ class TestTransformRequestBannedParams:
|
|||
)
|
||||
|
||||
|
||||
class TestTransformRequestOffEventLoop:
|
||||
@pytest.fixture
|
||||
def client(self):
|
||||
mock_auth = UserAPIKeyAuth(user_id="test-internal", user_role=LitellmUserRoles.INTERNAL_USER)
|
||||
original = app.dependency_overrides.copy()
|
||||
app.dependency_overrides[user_api_key_auth] = lambda: mock_auth
|
||||
try:
|
||||
yield TestClient(app)
|
||||
finally:
|
||||
app.dependency_overrides = original
|
||||
|
||||
def test_transform_request_runs_return_raw_request_off_the_event_loop(self, client, monkeypatch):
|
||||
import litellm.utils
|
||||
from litellm.types.utils import RawRequestTypedDict
|
||||
|
||||
seen: dict[str, bool] = {}
|
||||
|
||||
def fake_return_raw_request(endpoint, kwargs):
|
||||
try:
|
||||
asyncio.get_running_loop()
|
||||
seen["on_event_loop"] = True
|
||||
except RuntimeError:
|
||||
seen["on_event_loop"] = False
|
||||
return RawRequestTypedDict(
|
||||
raw_request_api_base="https://api.openai.com/v1/",
|
||||
raw_request_body=kwargs,
|
||||
raw_request_headers={},
|
||||
error=None,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(litellm.utils, "return_raw_request", fake_return_raw_request)
|
||||
response = client.post(
|
||||
"/utils/transform_request",
|
||||
json={
|
||||
"call_type": "completion",
|
||||
"request_body": {"model": "gpt-5.6-sol", "messages": [{"role": "user", "content": "hi"}]},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.json()["raw_request_body"]["model"] == "gpt-5.6-sol"
|
||||
assert seen == {"on_event_loop": False}, "return_raw_request ran on the event loop thread"
|
||||
|
||||
|
||||
class TestSortModelsByDisplayName:
|
||||
"""Regression: team BYOK rows persist an internal `model_name` like
|
||||
`model_name_{team_id}_{uuid}` and expose the user-facing name via
|
||||
|
|
@ -12531,9 +12612,7 @@ async def test_update_config_general_settings_refuses_a_key_the_config_file_decl
|
|||
admin = UserAPIKeyAuth(api_key="hashed-admin", user_id="admin-1", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
with pytest.raises(HTTPException) as excinfo:
|
||||
await update_config_general_settings(
|
||||
data=ConfigFieldUpdate(
|
||||
field_name="max_parallel_requests", field_value=999, config_type="general_settings"
|
||||
),
|
||||
data=ConfigFieldUpdate(field_name="max_parallel_requests", field_value=999, config_type="general_settings"),
|
||||
user_api_key_dict=admin,
|
||||
)
|
||||
|
||||
|
|
@ -14101,9 +14180,15 @@ async def test_moderations_response_carries_litellm_call_id_header():
|
|||
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", spend=0.0)
|
||||
|
||||
with (
|
||||
patch.object(proxy_server_module, "add_litellm_data_to_request", new=passthrough_add_litellm_data), # test-quality-ok: the route reads this module global, no injection point
|
||||
patch.object(proxy_server_module, "route_request", new=AsyncMock(return_value=fake_llm_call())), # test-quality-ok: fakes the provider call so the response headers assembled by the real route are observable
|
||||
patch.object(proxy_server_module, "proxy_logging_obj") as mock_logging, # test-quality-ok: module global, no injection point
|
||||
patch.object(
|
||||
proxy_server_module, "add_litellm_data_to_request", new=passthrough_add_litellm_data
|
||||
), # test-quality-ok: the route reads this module global, no injection point
|
||||
patch.object(
|
||||
proxy_server_module, "route_request", new=AsyncMock(return_value=fake_llm_call())
|
||||
), # test-quality-ok: fakes the provider call so the response headers assembled by the real route are observable
|
||||
patch.object(
|
||||
proxy_server_module, "proxy_logging_obj"
|
||||
) as mock_logging, # test-quality-ok: module global, no injection point
|
||||
):
|
||||
mock_logging.pre_call_hook = AsyncMock(side_effect=lambda user_api_key_dict, data, call_type: data)
|
||||
mock_logging.update_request_status = AsyncMock()
|
||||
|
|
@ -14140,9 +14225,15 @@ async def test_moderations_failure_log_carries_the_callers_litellm_call_id(caplo
|
|||
verbose_proxy_logger.propagate = True
|
||||
try:
|
||||
with (
|
||||
patch.object(proxy_server_module, "add_litellm_data_to_request", new=passthrough_add_litellm_data), # test-quality-ok: the route reads this module global, no injection point
|
||||
patch.object(proxy_server_module, "route_request", new=AsyncMock(side_effect=Exception("bad key"))), # test-quality-ok: fakes the provider failure so the real route's error log is observable
|
||||
patch.object(proxy_server_module, "proxy_logging_obj", new=fake_logging), # test-quality-ok: module global, no injection point
|
||||
patch.object(
|
||||
proxy_server_module, "add_litellm_data_to_request", new=passthrough_add_litellm_data
|
||||
), # test-quality-ok: the route reads this module global, no injection point
|
||||
patch.object(
|
||||
proxy_server_module, "route_request", new=AsyncMock(side_effect=Exception("bad key"))
|
||||
), # test-quality-ok: fakes the provider failure so the real route's error log is observable
|
||||
patch.object(
|
||||
proxy_server_module, "proxy_logging_obj", new=fake_logging
|
||||
), # test-quality-ok: module global, no injection point
|
||||
caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"),
|
||||
pytest.raises(ProxyException) as raised,
|
||||
):
|
||||
|
|
@ -14175,7 +14266,9 @@ async def test_moderations_unparseable_body_bills_the_callers_litellm_call_id():
|
|||
fake_logging.post_call_failure_hook = AsyncMock()
|
||||
|
||||
with (
|
||||
patch.object(proxy_server_module, "proxy_logging_obj", new=fake_logging), # test-quality-ok: module global, no injection point
|
||||
patch.object(
|
||||
proxy_server_module, "proxy_logging_obj", new=fake_logging
|
||||
), # test-quality-ok: module global, no injection point
|
||||
pytest.raises(ProxyException) as raised,
|
||||
):
|
||||
await proxy_server_module.moderations(
|
||||
|
|
@ -14203,8 +14296,12 @@ async def test_moderations_already_shaped_failure_answers_with_the_callers_litel
|
|||
fake_logging.post_call_failure_hook = AsyncMock()
|
||||
|
||||
with (
|
||||
patch.object(proxy_server_module, "add_litellm_data_to_request", new=AsyncMock(side_effect=exc)), # test-quality-ok: the route reads this module global, no injection point
|
||||
patch.object(proxy_server_module, "proxy_logging_obj", new=fake_logging), # test-quality-ok: module global, no injection point
|
||||
patch.object(
|
||||
proxy_server_module, "add_litellm_data_to_request", new=AsyncMock(side_effect=exc)
|
||||
), # test-quality-ok: the route reads this module global, no injection point
|
||||
patch.object(
|
||||
proxy_server_module, "proxy_logging_obj", new=fake_logging
|
||||
), # test-quality-ok: module global, no injection point
|
||||
pytest.raises(ProxyException) as raised,
|
||||
):
|
||||
await proxy_server_module.moderations(
|
||||
|
|
@ -14239,8 +14336,12 @@ async def test_audio_speech_already_shaped_failure_answers_with_the_callers_lite
|
|||
fake_logging.post_call_failure_hook = AsyncMock()
|
||||
|
||||
with (
|
||||
patch.object(proxy_server_module, "add_litellm_data_to_request", new=AsyncMock(side_effect=exc)), # test-quality-ok: the route reads this module global, no injection point
|
||||
patch.object(proxy_server_module, "proxy_logging_obj", new=fake_logging), # test-quality-ok: module global, no injection point
|
||||
patch.object(
|
||||
proxy_server_module, "add_litellm_data_to_request", new=AsyncMock(side_effect=exc)
|
||||
), # test-quality-ok: the route reads this module global, no injection point
|
||||
patch.object(
|
||||
proxy_server_module, "proxy_logging_obj", new=fake_logging
|
||||
), # test-quality-ok: module global, no injection point
|
||||
pytest.raises(type(exc)) as raised,
|
||||
):
|
||||
await proxy_server_module.audio_speech(
|
||||
|
|
@ -15025,14 +15126,18 @@ async def test_token_counter_loads_a_custom_tokenizer_off_the_event_loop(monkeyp
|
|||
{
|
||||
"model_name": "self-hosted",
|
||||
"litellm_params": {"model": "openai/self-hosted-model", "api_base": "http://localhost:8080/v1"},
|
||||
"model_info": {"custom_tokenizer": {"identifier": "my-org/tokenizer", "revision": "main", "auth_token": None}},
|
||||
"model_info": {
|
||||
"custom_tokenizer": {"identifier": "my-org/tokenizer", "revision": "main", "auth_token": None}
|
||||
},
|
||||
}
|
||||
]
|
||||
),
|
||||
)
|
||||
|
||||
response, took, lags = await timed_with_loop_lags(
|
||||
lambda: proxy_server_module.token_counter(TokenCountRequest(model="self-hosted", prompt="count me off the loop"))
|
||||
lambda: proxy_server_module.token_counter(
|
||||
TokenCountRequest(model="self-hosted", prompt="count me off the loop")
|
||||
)
|
||||
)
|
||||
|
||||
assert response.tokenizer_type == "huggingface_tokenizer"
|
||||
|
|
@ -15256,3 +15361,123 @@ async def test_initialize_jwt_auth_leaves_the_declared_jwtauth_mapping_unresolve
|
|||
|
||||
assert declared["team_id_jwt_field"] == "os.environ/JWT_TEAM_FIELD"
|
||||
assert proxy_server_module.jwt_handler.litellm_jwtauth.team_id_jwt_field == "resolved-team-field"
|
||||
|
||||
|
||||
def test_spend_capture_rate_check_job_validates_the_boot_settings_and_reads_them_again_on_every_run(monkeypatch):
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.constants import SPEND_CAPTURE_RATE_CHECK_JOB_ID
|
||||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||||
|
||||
scheduler = MagicMock()
|
||||
general_settings: dict[str, object] = {}
|
||||
seen_settings = []
|
||||
|
||||
async def fake_scheduled_check(prisma_client, settings, *, pod_lock_manager, alert, publish):
|
||||
seen_settings.append(settings)
|
||||
return ()
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.run_scheduled_spend_capture_rate_check", fake_scheduled_check)
|
||||
ProxyStartupEvent._initialize_spend_capture_rate_check_job(
|
||||
scheduler=scheduler,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
prisma_client=MagicMock(),
|
||||
read_general_settings=lambda: general_settings,
|
||||
)
|
||||
scheduler.add_job.assert_called_once()
|
||||
assert scheduler.add_job.call_args.kwargs["id"] == SPEND_CAPTURE_RATE_CHECK_JOB_ID
|
||||
check = scheduler.add_job.call_args.args[0]
|
||||
|
||||
asyncio.run(check())
|
||||
assert seen_settings == []
|
||||
|
||||
general_settings["spend_capture_rate_check"] = {"providers": ["openai"], "threshold": 0.85}
|
||||
asyncio.run(check())
|
||||
general_settings["spend_capture_rate_check"] = {"threshold": 0.7, "lookback_days": 3}
|
||||
asyncio.run(check())
|
||||
assert [(s.threshold, s.lookback_days) for s in seen_settings] == [(0.85, 7), (0.7, 3)]
|
||||
|
||||
with pytest.raises(ValidationError, match="threshhold"):
|
||||
ProxyStartupEvent._initialize_spend_capture_rate_check_job(
|
||||
scheduler=scheduler,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
prisma_client=MagicMock(),
|
||||
read_general_settings=lambda: {"spend_capture_rate_check": {"threshhold": 0.85}},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_spend_capture_rate_check_job_publishes_to_prometheus_and_alerts(monkeypatch):
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||||
|
||||
scheduler = MagicMock()
|
||||
proxy_logging = MagicMock()
|
||||
proxy_logging.alerting_handler = AsyncMock()
|
||||
proxy_logging.db_spend_update_writer.pod_lock_manager = None
|
||||
prometheus = MagicMock(spec=PrometheusLogger)
|
||||
monkeypatch.setattr(
|
||||
litellm.logging_callback_manager, "get_custom_loggers_for_type", lambda callback_type: [prometheus]
|
||||
)
|
||||
|
||||
async def fake_scheduled_check(prisma_client, settings, *, pod_lock_manager, alert, publish):
|
||||
publish("openai", 0.42)
|
||||
publish("openai", None)
|
||||
await alert("under the threshold")
|
||||
return ()
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.run_scheduled_spend_capture_rate_check", fake_scheduled_check)
|
||||
ProxyStartupEvent._initialize_spend_capture_rate_check_job(
|
||||
scheduler=scheduler,
|
||||
proxy_logging_obj=proxy_logging,
|
||||
prisma_client=MagicMock(),
|
||||
read_general_settings=lambda: {"spend_capture_rate_check": {}},
|
||||
)
|
||||
|
||||
await scheduler.add_job.call_args.args[0]()
|
||||
|
||||
assert prometheus.set_spend_capture_rate.call_args_list == [
|
||||
call(api_provider="openai", capture_rate=0.42),
|
||||
call(api_provider="openai", capture_rate=None),
|
||||
]
|
||||
proxy_logging.alerting_handler.assert_awaited_once()
|
||||
assert proxy_logging.alerting_handler.await_args.kwargs["message"] == "under the threshold"
|
||||
assert proxy_logging.alerting_handler.await_args.kwargs["level"] == "High"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_spend_capture_rate_check_job_clears_the_gauge_once_the_setting_is_removed(monkeypatch):
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||||
|
||||
scheduler = MagicMock()
|
||||
general_settings: dict[str, object] = {"spend_capture_rate_check": {}}
|
||||
prometheus = MagicMock(spec=PrometheusLogger)
|
||||
monkeypatch.setattr(
|
||||
litellm.logging_callback_manager, "get_custom_loggers_for_type", lambda callback_type: [prometheus]
|
||||
)
|
||||
scheduled_checks = []
|
||||
|
||||
async def fake_scheduled_check(prisma_client, settings, *, pod_lock_manager, alert, publish):
|
||||
scheduled_checks.append(settings)
|
||||
publish("openai", 0.97)
|
||||
return ()
|
||||
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.run_scheduled_spend_capture_rate_check", fake_scheduled_check)
|
||||
ProxyStartupEvent._initialize_spend_capture_rate_check_job(
|
||||
scheduler=scheduler,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
prisma_client=MagicMock(),
|
||||
read_general_settings=lambda: general_settings,
|
||||
)
|
||||
check = scheduler.add_job.call_args.args[0]
|
||||
|
||||
await check()
|
||||
del general_settings["spend_capture_rate_check"]
|
||||
await check()
|
||||
|
||||
assert len(scheduled_checks) == 1
|
||||
assert prometheus.set_spend_capture_rate.call_args_list == [
|
||||
call(api_provider="openai", capture_rate=0.97),
|
||||
call(api_provider="openai", capture_rate=None),
|
||||
]
|
||||
|
|
|
|||
|
|
@ -21,6 +21,10 @@ def _scoped_root(pathspec: str) -> str:
|
|||
return re.sub(r"^:\([^)]*\)", "", pathspec).split("*", 1)[0]
|
||||
|
||||
|
||||
def _gate_rooted_at(root: str) -> tuple[str, ...]:
|
||||
return next(gate for gate in GATES if _scoped_root(gate[0]) == root)
|
||||
|
||||
|
||||
def _changed_files_selected_by(tmp_path: Path, pathspecs: tuple[str, ...], files: tuple[str, ...]) -> frozenset[str]:
|
||||
_git(tmp_path, "init", "-q", "-b", "main")
|
||||
_git(tmp_path, "config", "user.email", "t@t")
|
||||
|
|
@ -37,8 +41,10 @@ def _changed_files_selected_by(tmp_path: Path, pathspecs: tuple[str, ...], files
|
|||
)
|
||||
|
||||
|
||||
def test_workflow_still_carries_the_ruff_format_and_e2e_basedpyright_diff_gates() -> None:
|
||||
assert frozenset(_scoped_root(gate[0]) for gate in GATES) == frozenset({"litellm/", "tests/e2e/"})
|
||||
def test_workflow_still_carries_the_ruff_format_e2e_basedpyright_and_claude_code_harness_diff_gates() -> None:
|
||||
assert frozenset(_scoped_root(gate[0]) for gate in GATES) == frozenset(
|
||||
{"litellm/", "tests/e2e/", "tests/e2e/claude_code/"}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("pathspecs", GATES, ids=" ".join)
|
||||
|
|
@ -52,3 +58,21 @@ def test_diff_gate_selects_top_level_and_nested_python_files_only(tmp_path: Path
|
|||
(top_level, nested, f"{root}notes.md", "elsewhere/top_level_module.py", "elsewhere/pkg/nested_module.py"),
|
||||
)
|
||||
assert selected == frozenset({top_level, nested})
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"trigger",
|
||||
(
|
||||
"tests/e2e/claude_code/cron_vm/install_claude_code.sh",
|
||||
"pyproject.toml",
|
||||
"uv.lock",
|
||||
".github/workflows/test-linting.yml",
|
||||
),
|
||||
)
|
||||
def test_claude_code_gate_also_fires_on_its_installer_dependency_manifests_and_workflow(
|
||||
tmp_path: Path, trigger: str
|
||||
) -> None:
|
||||
selected = _changed_files_selected_by(
|
||||
tmp_path, _gate_rooted_at("tests/e2e/claude_code/"), (trigger, "elsewhere/pyproject.toml", "tests/e2e/notes.md")
|
||||
)
|
||||
assert selected == frozenset({trigger})
|
||||
|
|
|
|||
|
|
@ -713,6 +713,37 @@ def _mocked_openai_chat_response(model: str) -> httpx.Response:
|
|||
)
|
||||
|
||||
|
||||
def test_return_raw_request_does_not_call_provider(respx_mock: respx.MockRouter):
|
||||
"""Regression for #33952: return_raw_request must transform without contacting the provider.
|
||||
|
||||
Previously return_raw_request invoked the real endpoint with a fake key and relied on the
|
||||
provider rejecting it, which sent an unintended inference request and (in the async proxy
|
||||
route) blocked the event loop on provider I/O.
|
||||
"""
|
||||
from litellm.types.utils import CallTypes
|
||||
from litellm.utils import return_raw_request
|
||||
|
||||
model = "gpt-4o"
|
||||
route = respx_mock.post("https://api.openai.com/v1/chat/completions").mock(
|
||||
return_value=_mocked_openai_chat_response(model)
|
||||
)
|
||||
|
||||
request = return_raw_request(
|
||||
endpoint=CallTypes.completion,
|
||||
kwargs={
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
},
|
||||
)
|
||||
|
||||
assert route.call_count == 0
|
||||
assert request.get("error") is None
|
||||
assert request["raw_request_body"]["model"] == model
|
||||
assert request["raw_request_body"]["messages"] == [
|
||||
{"role": "user", "content": "hi"}
|
||||
]
|
||||
|
||||
|
||||
def test_completion_forwards_verbosity_in_raw_request(respx_mock: respx.MockRouter):
|
||||
"""Regression test: completion() must forward the verbosity param to the provider request body."""
|
||||
from litellm.types.utils import CallTypes
|
||||
|
|
|
|||
|
|
@ -1672,6 +1672,49 @@ class TestBedrockBatchNonChatEndpointRecords:
|
|||
assert "input" not in model_input
|
||||
assert "max_output_tokens" not in model_input
|
||||
|
||||
def test_anthropic_responses_record_accepts_a_function_tool_without_strict(self):
|
||||
"""Clients omit the SDK's required `strict`; the record is forwarded like real time, not validated."""
|
||||
parameters = {"type": "object", "properties": {"city": {"type": "string"}}}
|
||||
model_input = self._transform(
|
||||
{
|
||||
"custom_id": "4a",
|
||||
"method": "POST",
|
||||
"url": "/v1/responses",
|
||||
"body": {
|
||||
"model": self.ANTHROPIC_MODEL,
|
||||
"input": "Weather in Paris?",
|
||||
"tools": [{"type": "function", "name": "get_weather", "parameters": parameters}],
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
assert model_input["messages"][0]["content"] == [{"type": "text", "text": "Weather in Paris?"}]
|
||||
tool = model_input["tools"][0]
|
||||
function = tool.get("function", tool)
|
||||
assert (function["name"], function.get("parameters", function.get("input_schema"))) == ("get_weather", parameters)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("url", "body"),
|
||||
[
|
||||
(
|
||||
"/v1/responses",
|
||||
{"input": [{"role": "developer", "content": "be terse"}, {"role": "user", "content": "ping"}]},
|
||||
),
|
||||
(
|
||||
"/v1/chat/completions",
|
||||
{"messages": [{"role": "developer", "content": "be terse"}, {"role": "user", "content": "ping"}]},
|
||||
),
|
||||
],
|
||||
ids=["responses", "chat"],
|
||||
)
|
||||
def test_anthropic_developer_role_becomes_the_system_prompt_like_real_time(self, url, body):
|
||||
model_input = self._transform(
|
||||
{"custom_id": "4c", "method": "POST", "url": url, "body": {"model": self.ANTHROPIC_MODEL, **body}}
|
||||
)
|
||||
|
||||
assert model_input["system"] == [{"type": "text", "text": "be terse"}]
|
||||
assert [message["role"] for message in model_input["messages"]] == ["user"]
|
||||
|
||||
def test_responses_record_keeps_metadata(self):
|
||||
"""`metadata` reaches the bridge, which reads it as its own kwarg."""
|
||||
model_input = self._transform(
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ Includes tests for Vertex AI batch output transformation to OpenAI format.
|
|||
|
||||
import json
|
||||
import urllib.parse
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
|
||||
|
|
@ -1447,6 +1448,158 @@ class TestVertexEmbeddingsBatchInputTranslation:
|
|||
assert "content" in embeddings_row["request"]
|
||||
|
||||
|
||||
def _responses_entry(
|
||||
body: Mapping[str, object] | None = None,
|
||||
custom_id: str = "resp-1",
|
||||
url: str = "/v1/responses",
|
||||
) -> dict[str, object]:
|
||||
return {
|
||||
"custom_id": custom_id,
|
||||
"method": "POST",
|
||||
"url": url,
|
||||
"body": body
|
||||
if body is not None
|
||||
else {"model": "gemini-2.5-flash", "input": "What was the top headline in world news yesterday?"},
|
||||
}
|
||||
|
||||
|
||||
class TestVertexResponsesBatchInputTranslation:
|
||||
"""
|
||||
/v1/responses batch lines carry `input`, not `messages`, so they go through the
|
||||
Responses-to-Chat bridge before the Gemini translation instead of uploading as an
|
||||
empty text part.
|
||||
"""
|
||||
|
||||
def test_string_input_becomes_the_user_prompt(self):
|
||||
(row,) = _wrap_entries([_responses_entry()])
|
||||
|
||||
assert row["request"]["contents"] == [
|
||||
{"role": "user", "parts": [{"text": "What was the top headline in world news yesterday?"}]}
|
||||
]
|
||||
assert row["request"]["labels"]["litellm_custom_id"] == "resp-1"
|
||||
|
||||
def test_instructions_and_input_items_map_like_real_time(self):
|
||||
(row,) = _wrap_entries(
|
||||
[
|
||||
_responses_entry(
|
||||
body={
|
||||
"model": "gemini-2.5-flash",
|
||||
"instructions": "be terse",
|
||||
"input": [
|
||||
{"role": "user", "content": "what is 2+2?"},
|
||||
{"role": "assistant", "content": "4"},
|
||||
{"role": "user", "content": "and 3+3?"},
|
||||
],
|
||||
"max_output_tokens": 32,
|
||||
"temperature": 0.2,
|
||||
}
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
request = row["request"]
|
||||
assert request["system_instruction"] == {"parts": [{"text": "be terse"}]}
|
||||
assert [content["role"] for content in request["contents"]] == ["user", "model", "user"]
|
||||
assert request["contents"][-1]["parts"] == [{"text": "and 3+3?"}]
|
||||
assert request["generationConfig"]["max_output_tokens"] == 32
|
||||
assert request["generationConfig"]["temperature"] == 0.2
|
||||
|
||||
def test_web_search_tool_keeps_the_prompt(self):
|
||||
(row,) = _wrap_entries(
|
||||
[
|
||||
_responses_entry(
|
||||
body={
|
||||
"model": "gemini-2.5-flash",
|
||||
"input": "What was the top headline in world news yesterday?",
|
||||
"tools": [{"type": "web_search"}],
|
||||
}
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
assert row["request"]["contents"] == [
|
||||
{"role": "user", "parts": [{"text": "What was the top headline in world news yesterday?"}]}
|
||||
]
|
||||
assert row["request"]["tools"]
|
||||
|
||||
def test_sdk_optional_keys_are_not_required_like_real_time(self):
|
||||
(row,) = _wrap_entries(
|
||||
[
|
||||
_responses_entry(
|
||||
body={
|
||||
"model": "gemini-2.5-flash",
|
||||
"input": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "input_text", "text": "Weather in the pictured city?"},
|
||||
{"type": "input_image", "image_url": "https://example.com/paris.png"},
|
||||
],
|
||||
}
|
||||
],
|
||||
"tools": [
|
||||
{
|
||||
"type": "function",
|
||||
"name": "get_weather",
|
||||
"parameters": {"type": "object", "properties": {"city": {"type": "string"}}},
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
request = row["request"]
|
||||
assert request["contents"][0]["parts"] == [
|
||||
{"text": "Weather in the pictured city?"},
|
||||
{"file_data": {"mime_type": "image/png", "file_uri": "https://example.com/paris.png"}},
|
||||
]
|
||||
assert request["tools"][0]["function_declarations"][0]["name"] == "get_weather"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"url",
|
||||
["/v1/responses", "/v1/responses/", "/v1/responses?beta=1", "responses", "https://api.openai.com/v1/responses"],
|
||||
)
|
||||
def test_route_spellings_are_all_responses(self, url):
|
||||
(row,) = _wrap_entries([_responses_entry(url=url)])
|
||||
|
||||
assert row["request"]["contents"][0]["parts"] == [
|
||||
{"text": "What was the top headline in world news yesterday?"}
|
||||
]
|
||||
|
||||
def test_missing_input_fails_the_upload(self):
|
||||
with pytest.raises(ValueError, match="missing required `input` field"):
|
||||
_wrap_entries([_responses_entry(body={"model": "gemini-2.5-flash"})])
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"entry",
|
||||
[
|
||||
_responses_entry(
|
||||
body={
|
||||
"model": "gemini-2.5-flash",
|
||||
"input": [{"role": "developer", "content": "be terse"}, {"role": "user", "content": "ping"}],
|
||||
}
|
||||
),
|
||||
{
|
||||
"custom_id": "chat-1",
|
||||
"method": "POST",
|
||||
"url": "/v1/chat/completions",
|
||||
"body": {
|
||||
"model": "gemini-2.5-flash",
|
||||
"messages": [{"role": "developer", "content": "be terse"}, {"role": "user", "content": "ping"}],
|
||||
},
|
||||
},
|
||||
],
|
||||
ids=["responses", "chat"],
|
||||
)
|
||||
def test_developer_role_becomes_the_system_instruction_like_real_time(self, entry):
|
||||
(row,) = _wrap_entries([entry])
|
||||
|
||||
request = row["request"]
|
||||
assert request["system_instruction"] == {"parts": [{"text": "be terse"}]}
|
||||
assert request["contents"] == [{"role": "user", "parts": [{"text": "ping"}]}]
|
||||
|
||||
|
||||
class TestVertexEmbeddingsBatchOutputTranslation:
|
||||
"""Vertex Gemini Embedding batch output rows must come back as OpenAI batch rows."""
|
||||
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@ import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/component
|
|||
import { Input } from "@/components/ui/input";
|
||||
import { Tooltip, TooltipContent, TooltipTrigger } from "@/components/ui/tooltip";
|
||||
import { Field, FieldDescription, FieldError, FieldGroup, FieldLabel } from "@/components/ui/field";
|
||||
import type { KeyValueFormValue, KillSwitchConfig, KillSwitchFormValue } from "./kill_switch_config";
|
||||
|
||||
export interface AgentSkillFormValue {
|
||||
id?: string;
|
||||
|
|
@ -54,6 +55,8 @@ export type AgentFormFieldValue =
|
|||
| string[]
|
||||
| AgentSkillFormValue[]
|
||||
| StaticHeaderFormValue[]
|
||||
| KeyValueFormValue[]
|
||||
| KillSwitchFormValue
|
||||
| McpServerSelection
|
||||
| Record<string, string[]>
|
||||
| null
|
||||
|
|
@ -82,6 +85,7 @@ export interface AgentFormValues {
|
|||
output_cost_per_token?: string | number;
|
||||
static_headers?: StaticHeaderFormValue[];
|
||||
extra_headers?: string[];
|
||||
kill_switch?: KillSwitchFormValue;
|
||||
tpm_limit?: number | null;
|
||||
rpm_limit?: number | null;
|
||||
session_tpm_limit?: number | null;
|
||||
|
|
@ -123,6 +127,7 @@ export interface AgentRequestPayload {
|
|||
litellm_params?: Record<string, unknown>;
|
||||
object_permission?: Record<string, unknown>;
|
||||
access_group_ids?: string[];
|
||||
kill_switch?: KillSwitchConfig | null;
|
||||
}
|
||||
|
||||
interface AgentFormFieldProps {
|
||||
|
|
|
|||
|
|
@ -0,0 +1,124 @@
|
|||
import React from "react";
|
||||
import { fireEvent, render, screen, waitFor, within } from "@testing-library/react";
|
||||
import { describe, it, expect, vi, beforeEach } from "vitest";
|
||||
import AgentKillSwitchDangerZone from "./AgentKillSwitchDangerZone";
|
||||
import * as networking from "@/components/networking";
|
||||
import { toast } from "@/lib/toast";
|
||||
|
||||
vi.mock("@/components/networking", () => ({
|
||||
triggerAgentKillSwitchCall: vi.fn(),
|
||||
}));
|
||||
|
||||
vi.mock("@/lib/toast", () => ({
|
||||
toast: { success: vi.fn(), error: vi.fn() },
|
||||
}));
|
||||
|
||||
const killSwitch = { url: "https://ops.example.com/kill", method: "DELETE" as const };
|
||||
|
||||
const renderZone = (props: Partial<React.ComponentProps<typeof AgentKillSwitchDangerZone>> = {}) =>
|
||||
render(
|
||||
<AgentKillSwitchDangerZone
|
||||
agentId="agent-1"
|
||||
agentName="support-agent"
|
||||
killSwitch={killSwitch}
|
||||
accessToken="sk-test"
|
||||
isAdmin={true}
|
||||
{...props}
|
||||
/>,
|
||||
);
|
||||
|
||||
const openDialog = () => {
|
||||
fireEvent.click(screen.getByRole("button", { name: "Fire Kill Switch" }));
|
||||
return screen.getByRole("dialog");
|
||||
};
|
||||
|
||||
const dialogFireButton = () => within(screen.getByRole("dialog")).getByRole("button", { name: "Fire Kill Switch" });
|
||||
|
||||
describe("AgentKillSwitchDangerZone", () => {
|
||||
beforeEach(() => {
|
||||
vi.mocked(networking.triggerAgentKillSwitchCall).mockReset();
|
||||
vi.mocked(toast.success).mockReset();
|
||||
vi.mocked(toast.error).mockReset();
|
||||
});
|
||||
|
||||
it("renders nothing for non-admins", () => {
|
||||
const { container } = renderZone({ isAdmin: false });
|
||||
|
||||
expect(container).toBeEmptyDOMElement();
|
||||
});
|
||||
|
||||
it("shows the webhook target and an outage warning inside a Danger Zone region", () => {
|
||||
renderZone();
|
||||
|
||||
const region = screen.getByRole("region", { name: "Danger Zone" });
|
||||
expect(region).toHaveTextContent("DELETE https://ops.example.com/kill");
|
||||
expect(region).toHaveTextContent("can cause an outage");
|
||||
expect(screen.getByRole("button", { name: "Fire Kill Switch" })).toBeEnabled();
|
||||
});
|
||||
|
||||
it("shows an unconfigured notice without a fire button when no kill switch is set", () => {
|
||||
renderZone({ killSwitch: null });
|
||||
|
||||
expect(screen.getByRole("region", { name: "Danger Zone" })).toHaveTextContent("Not configured");
|
||||
expect(screen.queryByRole("button", { name: "Fire Kill Switch" })).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("keeps the confirm button disabled until the exact agent name is typed", () => {
|
||||
renderZone();
|
||||
openDialog();
|
||||
|
||||
expect(dialogFireButton()).toBeDisabled();
|
||||
|
||||
fireEvent.change(screen.getByLabelText("Confirm agent name"), { target: { value: "support-agen" } });
|
||||
expect(dialogFireButton()).toBeDisabled();
|
||||
|
||||
fireEvent.change(screen.getByLabelText("Confirm agent name"), { target: { value: "support-agent" } });
|
||||
expect(dialogFireButton()).toBeEnabled();
|
||||
expect(networking.triggerAgentKillSwitchCall).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("fires the webhook after typed confirmation, closes the dialog and shows the sanitized result", async () => {
|
||||
const firedResult = {
|
||||
agent_id: "agent-1",
|
||||
url: killSwitch.url,
|
||||
method: "DELETE" as const,
|
||||
status_code: 202,
|
||||
response_body: '{"stopped": true}',
|
||||
};
|
||||
vi.mocked(networking.triggerAgentKillSwitchCall).mockResolvedValue(firedResult);
|
||||
renderZone();
|
||||
openDialog();
|
||||
|
||||
fireEvent.change(screen.getByLabelText("Confirm agent name"), { target: { value: "support-agent" } });
|
||||
fireEvent.click(dialogFireButton());
|
||||
|
||||
expect(await screen.findByRole("status")).toHaveTextContent('Last result: HTTP 202 {"stopped": true}');
|
||||
await waitFor(() => expect(screen.queryByRole("dialog")).not.toBeInTheDocument());
|
||||
expect(networking.triggerAgentKillSwitchCall).toHaveBeenCalledWith("sk-test", "agent-1");
|
||||
expect(toast.success).toHaveBeenCalledWith("Kill switch fired (HTTP 202)");
|
||||
});
|
||||
|
||||
it("does not call the webhook when the dialog is cancelled", () => {
|
||||
renderZone();
|
||||
openDialog();
|
||||
|
||||
fireEvent.change(screen.getByLabelText("Confirm agent name"), { target: { value: "support-agent" } });
|
||||
fireEvent.click(screen.getByRole("button", { name: "Cancel" }));
|
||||
|
||||
expect(networking.triggerAgentKillSwitchCall).not.toHaveBeenCalled();
|
||||
expect(screen.queryByRole("status")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("surfaces a failed webhook as an error toast and keeps the dialog open", async () => {
|
||||
vi.mocked(networking.triggerAgentKillSwitchCall).mockRejectedValue(new Error("Kill switch webhook returned 500"));
|
||||
renderZone();
|
||||
openDialog();
|
||||
|
||||
fireEvent.change(screen.getByLabelText("Confirm agent name"), { target: { value: "support-agent" } });
|
||||
fireEvent.click(dialogFireButton());
|
||||
|
||||
await waitFor(() => expect(toast.error).toHaveBeenCalledWith("Kill switch webhook returned 500"));
|
||||
expect(screen.getByRole("dialog")).toBeInTheDocument();
|
||||
expect(screen.queryByRole("status")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,147 @@
|
|||
import { CircleAlert } from "lucide-react";
|
||||
import React, { useState } from "react";
|
||||
import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog";
|
||||
import { InputGroup, InputGroupAddon, InputGroupInput } from "@/components/ui/input-group";
|
||||
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
|
||||
import { toast } from "@/lib/toast";
|
||||
import { AgentKillSwitchResult, triggerAgentKillSwitchCall } from "@/components/networking";
|
||||
import { KillSwitchConfig } from "./kill_switch_config";
|
||||
|
||||
interface AgentKillSwitchDangerZoneProps {
|
||||
agentId: string;
|
||||
agentName: string;
|
||||
killSwitch: KillSwitchConfig | null | undefined;
|
||||
accessToken: string | null;
|
||||
isAdmin: boolean;
|
||||
}
|
||||
|
||||
const AgentKillSwitchDangerZone: React.FC<AgentKillSwitchDangerZoneProps> = ({
|
||||
agentId,
|
||||
agentName,
|
||||
killSwitch,
|
||||
accessToken,
|
||||
isAdmin,
|
||||
}) => {
|
||||
const [isConfirmOpen, setIsConfirmOpen] = useState(false);
|
||||
const [confirmationInput, setConfirmationInput] = useState("");
|
||||
const [isFiring, setIsFiring] = useState(false);
|
||||
const [lastResult, setLastResult] = useState<AgentKillSwitchResult | null>(null);
|
||||
|
||||
if (!isAdmin) return null;
|
||||
|
||||
const openConfirm = () => {
|
||||
setConfirmationInput("");
|
||||
setIsConfirmOpen(true);
|
||||
};
|
||||
|
||||
const fire = async () => {
|
||||
if (!accessToken) return;
|
||||
setIsFiring(true);
|
||||
setLastResult(null);
|
||||
try {
|
||||
const result = await triggerAgentKillSwitchCall(accessToken, agentId);
|
||||
setLastResult(result);
|
||||
setIsConfirmOpen(false);
|
||||
toast.success(`Kill switch fired (HTTP ${result.status_code})`);
|
||||
} catch (error) {
|
||||
toast.error(error instanceof Error ? error.message : "Failed to fire kill switch");
|
||||
} finally {
|
||||
setIsFiring(false);
|
||||
}
|
||||
};
|
||||
|
||||
return (
|
||||
<section aria-labelledby="agent-danger-zone-heading" className="mt-6">
|
||||
<h3 id="agent-danger-zone-heading" className="text-lg font-medium text-destructive">
|
||||
Danger Zone
|
||||
</h3>
|
||||
<div className="mt-4 rounded-lg border border-destructive/40 bg-destructive/5 p-4">
|
||||
<div className="flex flex-wrap items-start justify-between gap-4">
|
||||
<div className="min-w-0 space-y-1 text-sm">
|
||||
<p className="font-medium text-foreground">Kill switch</p>
|
||||
{killSwitch ? (
|
||||
<>
|
||||
<p className="text-muted-foreground">
|
||||
Calls the configured webhook to stop this agent's upstream runtime. This can cause an outage for
|
||||
everyone using the agent and cannot be undone from LiteLLM
|
||||
</p>
|
||||
<p className="font-mono break-all text-foreground">
|
||||
{killSwitch.method ?? "POST"} {killSwitch.url}
|
||||
</p>
|
||||
</>
|
||||
) : (
|
||||
<p className="text-muted-foreground">
|
||||
Not configured. Add a kill switch webhook under Settings to enable this action
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
{killSwitch && (
|
||||
<Button type="button" variant="destructive" onClick={openConfirm} disabled={isFiring} aria-busy={isFiring}>
|
||||
{isFiring && <UiLoadingSpinner className="size-4" />}
|
||||
Fire Kill Switch
|
||||
</Button>
|
||||
)}
|
||||
</div>
|
||||
{lastResult && (
|
||||
<p className="mt-3 text-sm text-muted-foreground" role="status">
|
||||
Last result: HTTP {lastResult.status_code}
|
||||
{lastResult.response_body ? ` ${lastResult.response_body}` : ""}
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
|
||||
<Dialog open={isConfirmOpen} onOpenChange={(open) => !open && !isFiring && setIsConfirmOpen(false)}>
|
||||
<DialogContent className="max-h-[calc(100dvh-2rem)] overflow-y-auto">
|
||||
<DialogHeader>
|
||||
<DialogTitle>Fire kill switch for {agentName}?</DialogTitle>
|
||||
</DialogHeader>
|
||||
<div className="space-y-4">
|
||||
<Alert variant="error">
|
||||
<CircleAlert />
|
||||
<AlertTitle>This can cause an outage</AlertTitle>
|
||||
<AlertDescription>
|
||||
LiteLLM will call {killSwitch?.method ?? "POST"} {killSwitch?.url} immediately. Whatever that webhook
|
||||
does to the agent is outside LiteLLM's control and cannot be reverted here
|
||||
</AlertDescription>
|
||||
</Alert>
|
||||
<div>
|
||||
<p className="mb-2 text-base font-medium text-foreground">
|
||||
Type <span className="font-semibold text-destructive">{agentName}</span> to confirm:
|
||||
</p>
|
||||
<InputGroup className="rounded-md">
|
||||
<InputGroupAddon>
|
||||
<CircleAlert className="size-3.5 text-destructive" />
|
||||
</InputGroupAddon>
|
||||
<InputGroupInput
|
||||
aria-label="Confirm agent name"
|
||||
value={confirmationInput}
|
||||
onChange={(e) => setConfirmationInput(e.target.value)}
|
||||
placeholder={agentName}
|
||||
autoFocus
|
||||
/>
|
||||
</InputGroup>
|
||||
</div>
|
||||
</div>
|
||||
<DialogFooter>
|
||||
<Button variant="outline" onClick={() => setIsConfirmOpen(false)} disabled={isFiring}>
|
||||
Cancel
|
||||
</Button>
|
||||
<Button
|
||||
variant="destructive"
|
||||
onClick={fire}
|
||||
disabled={confirmationInput !== agentName || isFiring}
|
||||
aria-busy={isFiring}
|
||||
>
|
||||
{isFiring && <UiLoadingSpinner className="size-4" />}
|
||||
{isFiring ? "Firing..." : "Fire Kill Switch"}
|
||||
</Button>
|
||||
</DialogFooter>
|
||||
</DialogContent>
|
||||
</Dialog>
|
||||
</section>
|
||||
);
|
||||
};
|
||||
|
||||
export default AgentKillSwitchDangerZone;
|
||||
|
|
@ -0,0 +1,223 @@
|
|||
import React from "react";
|
||||
import { useFieldArray, useFormContext, useWatch } from "react-hook-form";
|
||||
import { Plus, Trash2 } from "lucide-react";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { Input } from "@/components/ui/input";
|
||||
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
|
||||
import { Textarea } from "@/components/ui/textarea";
|
||||
import { Field, FieldTitle } from "@/components/ui/field";
|
||||
import { PasswordInput } from "@/components/shared/PasswordInput";
|
||||
import { AgentFormField, AgentFormValues, labelWithHint } from "./AgentFormKit";
|
||||
import { KILL_SWITCH_AUTH_TYPES, KILL_SWITCH_METHODS, validateKillSwitchBody } from "./kill_switch_config";
|
||||
|
||||
const KeyValueFieldArray = ({
|
||||
name,
|
||||
addLabel,
|
||||
keyPlaceholder,
|
||||
valuePlaceholder,
|
||||
}: {
|
||||
name: "kill_switch.headers" | "kill_switch.query_params";
|
||||
addLabel: string;
|
||||
keyPlaceholder: string;
|
||||
valuePlaceholder: string;
|
||||
}) => {
|
||||
const { control } = useFormContext<AgentFormValues>();
|
||||
const { fields, append, remove } = useFieldArray({ control, name });
|
||||
|
||||
return (
|
||||
<div className="flex flex-col gap-2">
|
||||
{fields.map((item, index) => (
|
||||
<div key={item.id} className="flex items-start gap-2">
|
||||
<AgentFormField name={`${name}.${index}.key`} rules={{ required: "Name required" }}>
|
||||
{({ value, onChange, ref, ...control }) => (
|
||||
<Input
|
||||
{...control}
|
||||
ref={ref}
|
||||
className="w-55"
|
||||
placeholder={keyPlaceholder}
|
||||
value={typeof value === "string" ? value : ""}
|
||||
onChange={onChange}
|
||||
/>
|
||||
)}
|
||||
</AgentFormField>
|
||||
<AgentFormField name={`${name}.${index}.value`}>
|
||||
{({ value, onChange, ref, ...control }) => (
|
||||
<Input
|
||||
{...control}
|
||||
ref={ref}
|
||||
className="w-65"
|
||||
placeholder={valuePlaceholder}
|
||||
value={typeof value === "string" ? value : ""}
|
||||
onChange={onChange}
|
||||
/>
|
||||
)}
|
||||
</AgentFormField>
|
||||
<Button
|
||||
type="button"
|
||||
variant="ghost"
|
||||
size="icon"
|
||||
aria-label={`Remove ${addLabel.replace(/^Add /, "").toLowerCase()}`}
|
||||
className="text-destructive hover:text-destructive/80"
|
||||
onClick={() => remove(index)}
|
||||
>
|
||||
<Trash2 />
|
||||
</Button>
|
||||
</div>
|
||||
))}
|
||||
<Button type="button" variant="outline" className="w-full border-dashed" onClick={() => append({})}>
|
||||
<Plus />
|
||||
{addLabel}
|
||||
</Button>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
const TextField = ({
|
||||
name,
|
||||
label,
|
||||
placeholder,
|
||||
required,
|
||||
secret,
|
||||
}: {
|
||||
name: `kill_switch.${string}`;
|
||||
label: React.ReactNode;
|
||||
placeholder?: string;
|
||||
required?: string;
|
||||
secret?: boolean;
|
||||
}) => (
|
||||
<AgentFormField name={name} label={label} rules={required ? { required } : undefined}>
|
||||
{({ value, onChange, ref, ...control }) =>
|
||||
secret ? (
|
||||
<PasswordInput
|
||||
{...control}
|
||||
ref={ref}
|
||||
placeholder={placeholder}
|
||||
value={typeof value === "string" ? value : ""}
|
||||
onChange={onChange}
|
||||
/>
|
||||
) : (
|
||||
<Input
|
||||
{...control}
|
||||
ref={ref}
|
||||
placeholder={placeholder}
|
||||
value={typeof value === "string" ? value : ""}
|
||||
onChange={onChange}
|
||||
/>
|
||||
)
|
||||
}
|
||||
</AgentFormField>
|
||||
);
|
||||
|
||||
const KillSwitchAuthFields = () => {
|
||||
const { control } = useFormContext<AgentFormValues>();
|
||||
const authType = useWatch({ control, name: "kill_switch.auth_type" });
|
||||
|
||||
switch (authType) {
|
||||
case "bearer":
|
||||
return <TextField name="kill_switch.auth_token" label="Bearer token" required="Token required" secret />;
|
||||
case "api_key":
|
||||
return (
|
||||
<>
|
||||
<TextField name="kill_switch.auth_header_name" label="Header name" placeholder="X-API-Key" />
|
||||
<TextField name="kill_switch.auth_api_key" label="API key" required="API key required" secret />
|
||||
</>
|
||||
);
|
||||
case "basic":
|
||||
return (
|
||||
<>
|
||||
<TextField name="kill_switch.auth_username" label="Username" required="Username required" />
|
||||
<TextField name="kill_switch.auth_password" label="Password" required="Password required" secret />
|
||||
</>
|
||||
);
|
||||
default:
|
||||
return null;
|
||||
}
|
||||
};
|
||||
|
||||
const KillSwitchFormFields = () => (
|
||||
<>
|
||||
<TextField
|
||||
name="kill_switch.url"
|
||||
label={labelWithHint(
|
||||
"Webhook URL",
|
||||
"Absolute http(s) URL LiteLLM calls when the kill switch is triggered. Leave empty to remove the kill switch.",
|
||||
)}
|
||||
placeholder="https://example.com/hooks/kill-agent"
|
||||
/>
|
||||
|
||||
<AgentFormField name="kill_switch.method" label="Method">
|
||||
{({ value, onChange, ref: _ref, ...control }) => (
|
||||
<Select value={typeof value === "string" ? value : "POST"} onValueChange={onChange}>
|
||||
<SelectTrigger {...control} className="w-full">
|
||||
<SelectValue />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
{KILL_SWITCH_METHODS.map((method) => (
|
||||
<SelectItem key={method} value={method}>
|
||||
{method}
|
||||
</SelectItem>
|
||||
))}
|
||||
</SelectContent>
|
||||
</Select>
|
||||
)}
|
||||
</AgentFormField>
|
||||
|
||||
<Field>
|
||||
<FieldTitle>Headers</FieldTitle>
|
||||
<KeyValueFieldArray
|
||||
name="kill_switch.headers"
|
||||
addLabel="Add Header"
|
||||
keyPlaceholder="Header name"
|
||||
valuePlaceholder="Header value"
|
||||
/>
|
||||
</Field>
|
||||
|
||||
<Field>
|
||||
<FieldTitle>Query Parameters</FieldTitle>
|
||||
<KeyValueFieldArray
|
||||
name="kill_switch.query_params"
|
||||
addLabel="Add Query Parameter"
|
||||
keyPlaceholder="Parameter name"
|
||||
valuePlaceholder="Parameter value"
|
||||
/>
|
||||
</Field>
|
||||
|
||||
<AgentFormField
|
||||
name="kill_switch.body"
|
||||
label={labelWithHint("JSON Body", "Optional JSON object sent as the request body")}
|
||||
rules={{ validate: (value) => validateKillSwitchBody(typeof value === "string" ? value : "") }}
|
||||
>
|
||||
{({ value, onChange, ref, ...control }) => (
|
||||
<Textarea
|
||||
{...control}
|
||||
ref={ref}
|
||||
rows={4}
|
||||
placeholder='{"reason": "manual kill switch"}'
|
||||
value={typeof value === "string" ? value : ""}
|
||||
onChange={onChange}
|
||||
/>
|
||||
)}
|
||||
</AgentFormField>
|
||||
|
||||
<AgentFormField name="kill_switch.auth_type" label="Authentication">
|
||||
{({ value, onChange, ref: _ref, ...control }) => (
|
||||
<Select value={typeof value === "string" ? value : "none"} onValueChange={onChange}>
|
||||
<SelectTrigger {...control} className="w-full">
|
||||
<SelectValue />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
{KILL_SWITCH_AUTH_TYPES.map((option) => (
|
||||
<SelectItem key={option.value} value={option.value}>
|
||||
{option.label}
|
||||
</SelectItem>
|
||||
))}
|
||||
</SelectContent>
|
||||
</Select>
|
||||
)}
|
||||
</AgentFormField>
|
||||
|
||||
<KillSwitchAuthFields />
|
||||
</>
|
||||
);
|
||||
|
||||
export default KillSwitchFormFields;
|
||||
|
|
@ -3,6 +3,14 @@
|
|||
* Used across create, view, and update operations
|
||||
*/
|
||||
|
||||
import {
|
||||
EMPTY_KILL_SWITCH_FORM,
|
||||
buildKillSwitchFromForm,
|
||||
parseKillSwitchForForm,
|
||||
type KillSwitchConfig,
|
||||
type KillSwitchFormValue,
|
||||
} from "./kill_switch_config";
|
||||
|
||||
export interface FieldConfig {
|
||||
name: string;
|
||||
label: string;
|
||||
|
|
@ -236,6 +244,7 @@ export const getDefaultFormValues = () => {
|
|||
const defaults: any = {
|
||||
defaultInputModes: ["text"],
|
||||
defaultOutputModes: ["text"],
|
||||
kill_switch: { ...EMPTY_KILL_SWITCH_FORM },
|
||||
};
|
||||
|
||||
Object.values(AGENT_FORM_CONFIG).forEach((section) => {
|
||||
|
|
@ -310,9 +319,22 @@ export const buildAgentDataFromForm = (values: any, existingAgent?: any) => {
|
|||
agentData.extra_headers = values.extra_headers;
|
||||
}
|
||||
|
||||
applyKillSwitchToPayload(agentData, values.kill_switch, existingAgent);
|
||||
|
||||
return agentData;
|
||||
};
|
||||
|
||||
export const applyKillSwitchToPayload = (
|
||||
agentData: { kill_switch?: KillSwitchConfig | null },
|
||||
form: KillSwitchFormValue | undefined,
|
||||
existingAgent?: { kill_switch?: KillSwitchConfig | null },
|
||||
) => {
|
||||
const killSwitch = buildKillSwitchFromForm(form);
|
||||
if (killSwitch !== undefined && (killSwitch !== null || existingAgent?.kill_switch)) {
|
||||
agentData.kill_switch = killSwitch;
|
||||
}
|
||||
};
|
||||
|
||||
export const parseAccessGroupIdsForForm = (agent: { access_group_ids?: string[] | null }) => ({
|
||||
access_group_ids: agent.access_group_ids ?? [],
|
||||
});
|
||||
|
|
@ -380,6 +402,7 @@ export const parseAgentForForm = (agent: any) => {
|
|||
: [],
|
||||
// extra_headers: already an array of strings
|
||||
extra_headers: agent.extra_headers ?? [],
|
||||
kill_switch: parseKillSwitchForForm(agent.kill_switch),
|
||||
...parseMcpPermissionsForForm(agent),
|
||||
...parseAccessGroupIdsForForm(agent),
|
||||
};
|
||||
|
|
|
|||
|
|
@ -9,6 +9,8 @@ import { Textarea } from "@/components/ui/textarea";
|
|||
import { Field, FieldGroup, FieldTitle } from "@/components/ui/field";
|
||||
import { AGENT_FORM_CONFIG, SKILL_FIELD_CONFIG } from "./agent_config";
|
||||
import CostConfigFields, { COST_FIELD_NAMES } from "./cost_config_fields";
|
||||
import KillSwitchFormFields from "./KillSwitchFormFields";
|
||||
import { KILL_SWITCH_PANEL_KEY } from "./kill_switch_config";
|
||||
import {
|
||||
AgentFormField,
|
||||
AgentFormPanel,
|
||||
|
|
@ -30,6 +32,7 @@ export const A2A_PANEL_FIELD_NAMES: Readonly<Record<string, readonly string[]>>
|
|||
[AGENT_FORM_CONFIG.cost.key]: COST_FIELD_NAMES,
|
||||
[AGENT_FORM_CONFIG.litellm.key]: namesOf(AGENT_FORM_CONFIG.litellm.fields),
|
||||
[AUTH_HEADERS_PANEL_KEY]: ["static_headers", "extra_headers"],
|
||||
[KILL_SWITCH_PANEL_KEY]: ["kill_switch"],
|
||||
};
|
||||
|
||||
export const unmountedA2AFieldNames = (mountedPanels: readonly string[]): readonly string[] =>
|
||||
|
|
@ -394,6 +397,12 @@ const AgentFormFields: React.FC<AgentFormFieldsProps> = ({ panels, showAgentName
|
|||
</AgentFormField>
|
||||
</AgentFormPanel>
|
||||
)}
|
||||
|
||||
{shouldShow(KILL_SWITCH_PANEL_KEY) && (
|
||||
<AgentFormPanel panelKey={KILL_SWITCH_PANEL_KEY} title="Kill Switch" panels={panels}>
|
||||
<KillSwitchFormFields />
|
||||
</AgentFormPanel>
|
||||
)}
|
||||
</div>
|
||||
</>
|
||||
);
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ vi.mock("@/components/networking", () => ({
|
|||
getAgentInfo: vi.fn(),
|
||||
getAgentCreateMetadata: vi.fn(),
|
||||
patchAgentCall: vi.fn(),
|
||||
triggerAgentKillSwitchCall: vi.fn(),
|
||||
}));
|
||||
|
||||
vi.mock("@/app/(dashboard)/hooks/keys/useKeys", () => ({
|
||||
|
|
@ -75,9 +76,11 @@ const agent = {
|
|||
|
||||
describe("AgentInfoView settings", () => {
|
||||
beforeEach(() => {
|
||||
vi.restoreAllMocks();
|
||||
vi.mocked(networking.getAgentInfo).mockReset().mockResolvedValue(agent);
|
||||
vi.mocked(networking.getAgentCreateMetadata).mockReset().mockResolvedValue([]);
|
||||
vi.mocked(networking.patchAgentCall).mockReset().mockResolvedValue({});
|
||||
vi.mocked(networking.triggerAgentKillSwitchCall).mockReset();
|
||||
});
|
||||
|
||||
it("submits the edited agent when Save Changes is pressed", async () => {
|
||||
|
|
@ -158,4 +161,29 @@ describe("AgentInfoView settings", () => {
|
|||
expect(await screen.findByText("Access Groups")).toBeInTheDocument();
|
||||
expect(screen.getByText("None")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("renders the kill switch Danger Zone for admins with the configured webhook", async () => {
|
||||
vi.mocked(networking.getAgentInfo).mockResolvedValue({
|
||||
...agent,
|
||||
kill_switch: { url: "https://ops.example.com/kill", method: "DELETE" },
|
||||
});
|
||||
render(<AgentInfoView agentId="agent-1" onClose={vi.fn()} accessToken="sk-test" isAdmin={true} />);
|
||||
|
||||
const dangerZone = await screen.findByRole("region", { name: "Danger Zone" });
|
||||
expect(dangerZone).toHaveTextContent("DELETE https://ops.example.com/kill");
|
||||
expect(screen.getByRole("button", { name: "Fire Kill Switch" })).toBeInTheDocument();
|
||||
expect(screen.queryByText("Kill Switch")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("hides the Danger Zone from non-admins", async () => {
|
||||
vi.mocked(networking.getAgentInfo).mockResolvedValue({
|
||||
...agent,
|
||||
kill_switch: { url: "https://ops.example.com/kill", method: "POST" },
|
||||
});
|
||||
render(<AgentInfoView agentId="agent-1" onClose={vi.fn()} accessToken="sk-test" isAdmin={false} />);
|
||||
|
||||
expect(await screen.findByRole("heading", { name: "support-agent" })).toBeInTheDocument();
|
||||
expect(screen.queryByRole("region", { name: "Danger Zone" })).not.toBeInTheDocument();
|
||||
expect(screen.queryByRole("button", { name: "Fire Kill Switch" })).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ import KeyInfoView from "@/components/templates/key_info_view";
|
|||
import MCPServerSelector from "@/components/mcp_server_management/MCPServerSelector";
|
||||
import MCPToolPermissions from "@/components/mcp_server_management/MCPToolPermissions";
|
||||
import AgentVirtualKeys from "./agent_virtual_keys";
|
||||
import AgentKillSwitchDangerZone from "./AgentKillSwitchDangerZone";
|
||||
import AgentFormFields, { unmountedA2AFieldNames } from "./agent_form_fields";
|
||||
import DynamicAgentFormFields, { buildDynamicAgentData, unmountedDynamicFieldNames } from "./dynamic_agent_form_fields";
|
||||
import {
|
||||
|
|
@ -228,7 +229,7 @@ const AgentInfoView: React.FC<AgentInfoViewProps> = ({ agentId, onClose, accessT
|
|||
);
|
||||
|
||||
const built: AgentRequestPayload = usesDynamicFields
|
||||
? { ...buildDynamicAgentData(values, selectedAgentTypeInfo), agent_name: values.agent_name }
|
||||
? { ...buildDynamicAgentData(values, selectedAgentTypeInfo, agent), agent_name: values.agent_name }
|
||||
: buildAgentDataFromForm(values, agent);
|
||||
|
||||
const updateData = appliedDiscoveredSelection
|
||||
|
|
@ -459,6 +460,14 @@ const AgentInfoView: React.FC<AgentInfoViewProps> = ({ agentId, onClose, accessT
|
|||
</DetailList>
|
||||
</div>
|
||||
)}
|
||||
|
||||
<AgentKillSwitchDangerZone
|
||||
agentId={agent.agent_id}
|
||||
agentName={agent.agent_name}
|
||||
killSwitch={agent.kill_switch}
|
||||
accessToken={accessToken}
|
||||
isAdmin={isAdmin}
|
||||
/>
|
||||
</TabsContent>
|
||||
|
||||
{/* Settings Panel (only for admins) */}
|
||||
|
|
|
|||
|
|
@ -58,6 +58,32 @@ describe("parseDynamicAgentForForm", () => {
|
|||
|
||||
expect(values.agent_runtime_arn).toBe(FULL_RUNTIME_ARN);
|
||||
});
|
||||
|
||||
it("loads the stored kill switch into the edit form for non-A2A agents", () => {
|
||||
const agent = {
|
||||
agent_id: "agent-1",
|
||||
agent_name: "bedrock-agent",
|
||||
agent_card_params: { description: "" },
|
||||
litellm_params: { model: `bedrock/agentcore/${FULL_RUNTIME_ARN}` },
|
||||
kill_switch: {
|
||||
url: "https://ops.example.com/kill",
|
||||
method: "DELETE",
|
||||
headers: { "X-Env": "prod" },
|
||||
auth: { type: "bearer", token: "REDACTED_BY_LITELM" },
|
||||
},
|
||||
} as unknown as Agent;
|
||||
|
||||
const values = parseDynamicAgentForForm(agent, bedrockAgentcoreInfo);
|
||||
|
||||
const expectedForm = {
|
||||
url: "https://ops.example.com/kill",
|
||||
method: "DELETE",
|
||||
headers: [{ key: "X-Env", value: "prod" }],
|
||||
auth_type: "bearer",
|
||||
auth_token: "REDACTED_BY_LITELM",
|
||||
};
|
||||
expect(values.kill_switch).toMatchObject(expectedForm);
|
||||
});
|
||||
});
|
||||
|
||||
describe("detectAgentType", () => {
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import { Agent } from "@/components/agents/types";
|
||||
import { AgentCreateInfo } from "@/components/networking";
|
||||
import { parseKillSwitchForForm } from "./kill_switch_config";
|
||||
|
||||
/**
|
||||
* Detects the agent type from an agent's litellm_params.
|
||||
|
|
@ -77,6 +78,7 @@ export const parseDynamicAgentForForm = (agent: Agent, agentTypeInfo: AgentCreat
|
|||
values.cost_per_query = agent.litellm_params?.cost_per_query;
|
||||
values.input_cost_per_token = agent.litellm_params?.input_cost_per_token;
|
||||
values.output_cost_per_token = agent.litellm_params?.output_cost_per_token;
|
||||
values.kill_switch = parseKillSwitchForForm(agent.kill_switch);
|
||||
|
||||
return values;
|
||||
};
|
||||
|
|
|
|||
|
|
@ -0,0 +1,65 @@
|
|||
import { describe, expect, it } from "vitest";
|
||||
import type { AgentCreateInfo } from "@/components/networking";
|
||||
import type { AgentFormValues } from "./AgentFormKit";
|
||||
import { AGENT_FORM_CONFIG } from "./agent_config";
|
||||
import { buildDynamicAgentData, unmountedDynamicFieldNames } from "./dynamic_agent_form_fields";
|
||||
import {
|
||||
EMPTY_KILL_SWITCH_FORM,
|
||||
KILL_SWITCH_PANEL_KEY,
|
||||
type KillSwitchConfig,
|
||||
type KillSwitchFormValue,
|
||||
} from "./kill_switch_config";
|
||||
|
||||
const langgraphInfo: AgentCreateInfo = {
|
||||
agent_type: "langgraph",
|
||||
agent_type_display_name: "LangGraph",
|
||||
model_template: "langgraph/{assistant_id}",
|
||||
credential_fields: [{ key: "assistant_id", label: "Assistant ID", required: true, include_in_litellm_params: false }],
|
||||
};
|
||||
|
||||
const killSwitchForm: KillSwitchFormValue = {
|
||||
...EMPTY_KILL_SWITCH_FORM,
|
||||
url: "https://ops.example.com/kill",
|
||||
method: "DELETE",
|
||||
headers: [{ key: "X-Env", value: "prod" }],
|
||||
auth_type: "bearer",
|
||||
auth_token: "tok",
|
||||
};
|
||||
|
||||
const baseValues: AgentFormValues = { agent_name: "lg-agent", assistant_id: "asst_1" };
|
||||
|
||||
describe("buildDynamicAgentData kill switch", () => {
|
||||
it("serializes the kill switch section into the payload", () => {
|
||||
const payload = buildDynamicAgentData({ ...baseValues, kill_switch: killSwitchForm }, langgraphInfo);
|
||||
|
||||
const expected: KillSwitchConfig = {
|
||||
url: "https://ops.example.com/kill",
|
||||
method: "DELETE",
|
||||
headers: { "X-Env": "prod" },
|
||||
query_params: {},
|
||||
body: null,
|
||||
auth: { type: "bearer", token: "tok" },
|
||||
};
|
||||
expect(payload.kill_switch).toEqual(expected);
|
||||
expect(payload.litellm_params).toMatchObject({ model: "langgraph/asst_1" });
|
||||
});
|
||||
|
||||
it("leaves kill_switch off the payload when the section was never mounted", () => {
|
||||
expect("kill_switch" in buildDynamicAgentData(baseValues, langgraphInfo)).toBe(false);
|
||||
});
|
||||
|
||||
it("clears a stored kill switch when the URL is blanked, but not on a fresh create", () => {
|
||||
const blanked: AgentFormValues = { ...baseValues, kill_switch: { ...EMPTY_KILL_SWITCH_FORM } };
|
||||
|
||||
expect("kill_switch" in buildDynamicAgentData(blanked, langgraphInfo)).toBe(false);
|
||||
const stored = { kill_switch: { url: "https://old.example", method: "POST" as const } };
|
||||
expect(buildDynamicAgentData(blanked, langgraphInfo, stored).kill_switch).toBeNull();
|
||||
});
|
||||
});
|
||||
|
||||
describe("unmountedDynamicFieldNames", () => {
|
||||
it("drops kill_switch from the submit only while its panel is unmounted", () => {
|
||||
expect(unmountedDynamicFieldNames([AGENT_FORM_CONFIG.cost.key])).toEqual(["kill_switch"]);
|
||||
expect(unmountedDynamicFieldNames([AGENT_FORM_CONFIG.cost.key, KILL_SWITCH_PANEL_KEY])).toEqual([]);
|
||||
});
|
||||
});
|
||||
|
|
@ -5,8 +5,10 @@ import { Textarea } from "@/components/ui/textarea";
|
|||
import { FieldGroup } from "@/components/ui/field";
|
||||
import { AgentCreateInfo, AgentCredentialFieldMetadata } from "@/components/networking";
|
||||
import { PasswordInput } from "@/components/shared/PasswordInput";
|
||||
import { AGENT_FORM_CONFIG } from "./agent_config";
|
||||
import { AGENT_FORM_CONFIG, applyKillSwitchToPayload } from "./agent_config";
|
||||
import CostConfigFields, { COST_FIELD_NAMES } from "./cost_config_fields";
|
||||
import KillSwitchFormFields from "./KillSwitchFormFields";
|
||||
import { KILL_SWITCH_PANEL_KEY, type KillSwitchConfig } from "./kill_switch_config";
|
||||
import {
|
||||
AgentFormField,
|
||||
AgentFormPanel,
|
||||
|
|
@ -21,8 +23,10 @@ interface DynamicAgentFormFieldsProps {
|
|||
panels: CollapsiblePanelsState;
|
||||
}
|
||||
|
||||
export const unmountedDynamicFieldNames = (mountedPanels: readonly string[]): readonly string[] =>
|
||||
mountedPanels.includes(AGENT_FORM_CONFIG.cost.key) ? [] : COST_FIELD_NAMES;
|
||||
export const unmountedDynamicFieldNames = (mountedPanels: readonly string[]): readonly string[] => [
|
||||
...(mountedPanels.includes(AGENT_FORM_CONFIG.cost.key) ? [] : COST_FIELD_NAMES),
|
||||
...(mountedPanels.includes(KILL_SWITCH_PANEL_KEY) ? [] : ["kill_switch"]),
|
||||
];
|
||||
|
||||
// A field's validation_pattern is server-supplied metadata; if it's ever not a valid regex, skip
|
||||
// validation rather than throwing during render and taking the whole form down with it.
|
||||
|
|
@ -143,11 +147,18 @@ const DynamicAgentFormFields: React.FC<DynamicAgentFormFieldsProps> = ({ agentTy
|
|||
<AgentFormPanel panelKey={AGENT_FORM_CONFIG.cost.key} title={AGENT_FORM_CONFIG.cost.title} panels={panels}>
|
||||
<CostConfigFields />
|
||||
</AgentFormPanel>
|
||||
<AgentFormPanel panelKey={KILL_SWITCH_PANEL_KEY} title="Kill Switch" panels={panels}>
|
||||
<KillSwitchFormFields />
|
||||
</AgentFormPanel>
|
||||
</div>
|
||||
</>
|
||||
);
|
||||
|
||||
export const buildDynamicAgentData = (values: AgentFormValues, agentTypeInfo: AgentCreateInfo): AgentRequestPayload => {
|
||||
export const buildDynamicAgentData = (
|
||||
values: AgentFormValues,
|
||||
agentTypeInfo: AgentCreateInfo,
|
||||
existingAgent?: { kill_switch?: KillSwitchConfig | null },
|
||||
): AgentRequestPayload => {
|
||||
const litellmParams: Record<string, unknown> = {
|
||||
...(agentTypeInfo.litellm_params_template || {}),
|
||||
};
|
||||
|
|
@ -207,6 +218,8 @@ export const buildDynamicAgentData = (values: AgentFormValues, agentTypeInfo: Ag
|
|||
if (values.session_tpm_limit != null) agentData.session_tpm_limit = values.session_tpm_limit;
|
||||
if (values.session_rpm_limit != null) agentData.session_rpm_limit = values.session_rpm_limit;
|
||||
|
||||
applyKillSwitchToPayload(agentData, values.kill_switch, existingAgent);
|
||||
|
||||
return agentData;
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,118 @@
|
|||
import { describe, expect, it } from "vitest";
|
||||
import {
|
||||
EMPTY_KILL_SWITCH_FORM,
|
||||
buildKillSwitchFromForm,
|
||||
parseKillSwitchForForm,
|
||||
validateKillSwitchBody,
|
||||
type KillSwitchConfig,
|
||||
type KillSwitchFormValue,
|
||||
} from "./kill_switch_config";
|
||||
|
||||
const fullConfig: KillSwitchConfig = {
|
||||
url: "https://ops.example.com/kill?env=prod",
|
||||
method: "DELETE",
|
||||
headers: { "X-Env": "prod" },
|
||||
query_params: { agent: "billing-bot" },
|
||||
body: { reason: "manual stop", force: true },
|
||||
auth: { type: "api_key", header_name: "X-Ops-Key", api_key: "k-456" },
|
||||
};
|
||||
|
||||
describe("buildKillSwitchFromForm", () => {
|
||||
it("returns undefined when the form never touched the kill switch", () => {
|
||||
expect(buildKillSwitchFromForm(undefined)).toBeUndefined();
|
||||
});
|
||||
|
||||
it("returns null when the URL is blank so the backend clears the config", () => {
|
||||
const blankUrlForm: KillSwitchFormValue = {
|
||||
...EMPTY_KILL_SWITCH_FORM,
|
||||
url: " ",
|
||||
auth_type: "bearer",
|
||||
auth_token: "t",
|
||||
};
|
||||
expect(buildKillSwitchFromForm(blankUrlForm)).toBeNull();
|
||||
});
|
||||
|
||||
it("builds the full config, dropping rows without a key and parsing the JSON body", () => {
|
||||
const fullForm: KillSwitchFormValue = {
|
||||
url: " https://ops.example.com/kill?env=prod ",
|
||||
method: "DELETE",
|
||||
headers: [
|
||||
{ key: "X-Env", value: "prod" },
|
||||
{ key: " ", value: "ignored" },
|
||||
],
|
||||
query_params: [{ key: "agent", value: "billing-bot" }],
|
||||
body: '{"reason": "manual stop", "force": true}',
|
||||
auth_type: "api_key",
|
||||
auth_header_name: "X-Ops-Key",
|
||||
auth_api_key: "k-456",
|
||||
};
|
||||
expect(buildKillSwitchFromForm(fullForm)).toEqual(fullConfig);
|
||||
});
|
||||
|
||||
it.each([
|
||||
[{ auth_type: "none" as const, auth_token: "leftover" }, null],
|
||||
[
|
||||
{ auth_type: "bearer" as const, auth_token: "tok-123" },
|
||||
{ type: "bearer", token: "tok-123" },
|
||||
],
|
||||
[
|
||||
{ auth_type: "api_key" as const, auth_api_key: "k" },
|
||||
{ type: "api_key", header_name: "X-API-Key", api_key: "k" },
|
||||
],
|
||||
[
|
||||
{ auth_type: "basic" as const, auth_username: "ops", auth_password: "pw" },
|
||||
{ type: "basic", username: "ops", password: "pw" },
|
||||
],
|
||||
])("maps auth form fields %j to %j", (authFields, expectedAuth) => {
|
||||
const authForm: KillSwitchFormValue = { ...EMPTY_KILL_SWITCH_FORM, url: "https://x.example", ...authFields };
|
||||
expect(buildKillSwitchFromForm(authForm)?.auth).toEqual(expectedAuth);
|
||||
});
|
||||
|
||||
it("sends an empty body as null and defaults the method to POST", () => {
|
||||
const expected: KillSwitchConfig = {
|
||||
url: "https://x.example",
|
||||
method: "POST",
|
||||
headers: {},
|
||||
query_params: {},
|
||||
body: null,
|
||||
auth: null,
|
||||
};
|
||||
expect(buildKillSwitchFromForm({ url: "https://x.example", body: " " })).toEqual(expected);
|
||||
});
|
||||
});
|
||||
|
||||
describe("validateKillSwitchBody", () => {
|
||||
it.each(["", " ", undefined, '{"a": 1}'])("accepts %j", (text) => {
|
||||
expect(validateKillSwitchBody(text)).toBe(true);
|
||||
});
|
||||
|
||||
it.each(["[1, 2]", '"text"', "42", "null"])("rejects non-object JSON %s", (text) => {
|
||||
expect(validateKillSwitchBody(text)).toBe("Body must be a JSON object");
|
||||
});
|
||||
|
||||
it("rejects malformed JSON with the parser message", () => {
|
||||
expect(validateKillSwitchBody("{not json")).toMatch(/JSON/);
|
||||
});
|
||||
});
|
||||
|
||||
describe("parseKillSwitchForForm", () => {
|
||||
it("returns the empty form for a missing config", () => {
|
||||
expect(parseKillSwitchForForm(null)).toEqual(EMPTY_KILL_SWITCH_FORM);
|
||||
expect(parseKillSwitchForForm(undefined)).toEqual(EMPTY_KILL_SWITCH_FORM);
|
||||
});
|
||||
|
||||
it("round-trips a full config through the form representation", () => {
|
||||
expect(buildKillSwitchFromForm(parseKillSwitchForForm(fullConfig))).toEqual(fullConfig);
|
||||
});
|
||||
|
||||
it("keeps the redacted secret marker in the auth field so the backend restores it", () => {
|
||||
const parsed = parseKillSwitchForForm({
|
||||
url: "https://x.example",
|
||||
method: "POST",
|
||||
auth: { type: "bearer", token: "redacted-marker" },
|
||||
});
|
||||
expect(parsed.auth_type).toBe("bearer");
|
||||
expect(parsed.auth_token).toBe("redacted-marker");
|
||||
expect(parsed.auth_password).toBe("");
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,124 @@
|
|||
import type { components } from "@/lib/http/schema";
|
||||
|
||||
export type KillSwitchConfig = components["schemas"]["AgentKillSwitchConfig"];
|
||||
export type KillSwitchAuth = NonNullable<KillSwitchConfig["auth"]>;
|
||||
export type KillSwitchMethod = NonNullable<KillSwitchConfig["method"]>;
|
||||
export type KillSwitchAuthType = KillSwitchAuth["type"] | "none";
|
||||
|
||||
export const KILL_SWITCH_METHODS: readonly KillSwitchMethod[] = ["POST", "PUT", "PATCH", "DELETE", "GET"];
|
||||
export const KILL_SWITCH_AUTH_TYPES: readonly { value: KillSwitchAuthType; label: string }[] = [
|
||||
{ value: "none", label: "None" },
|
||||
{ value: "bearer", label: "Bearer token" },
|
||||
{ value: "api_key", label: "API key header" },
|
||||
{ value: "basic", label: "Basic auth" },
|
||||
];
|
||||
|
||||
export interface KeyValueFormValue {
|
||||
key?: string;
|
||||
value?: string;
|
||||
}
|
||||
|
||||
export interface KillSwitchFormValue {
|
||||
url?: string;
|
||||
method?: KillSwitchMethod;
|
||||
headers?: KeyValueFormValue[];
|
||||
query_params?: KeyValueFormValue[];
|
||||
body?: string;
|
||||
auth_type?: KillSwitchAuthType;
|
||||
auth_token?: string;
|
||||
auth_header_name?: string;
|
||||
auth_api_key?: string;
|
||||
auth_username?: string;
|
||||
auth_password?: string;
|
||||
}
|
||||
|
||||
export const KILL_SWITCH_PANEL_KEY = "kill_switch";
|
||||
|
||||
export const EMPTY_KILL_SWITCH_FORM: Readonly<KillSwitchFormValue> = {
|
||||
url: "",
|
||||
method: "POST",
|
||||
headers: [],
|
||||
query_params: [],
|
||||
body: "",
|
||||
auth_type: "none",
|
||||
};
|
||||
|
||||
const pairsToRecord = (pairs: readonly KeyValueFormValue[] | undefined): Record<string, string> =>
|
||||
Object.fromEntries(
|
||||
(pairs ?? []).map((pair) => [pair.key?.trim() ?? "", pair.value ?? ""] as const).filter(([key]) => key.length > 0),
|
||||
);
|
||||
|
||||
const recordToPairs = (record: Record<string, string> | undefined | null): KeyValueFormValue[] =>
|
||||
Object.entries(record ?? {}).map(([key, value]) => ({ key, value }));
|
||||
|
||||
export const parseKillSwitchBody = (text: string | undefined): Record<string, unknown> | null => {
|
||||
const trimmed = text?.trim() ?? "";
|
||||
if (trimmed.length === 0) return null;
|
||||
const parsed: unknown = JSON.parse(trimmed);
|
||||
if (parsed === null || typeof parsed !== "object" || Array.isArray(parsed)) {
|
||||
throw new Error("Body must be a JSON object");
|
||||
}
|
||||
return parsed as Record<string, unknown>;
|
||||
};
|
||||
|
||||
export const validateKillSwitchBody = (text: string | undefined): true | string => {
|
||||
try {
|
||||
parseKillSwitchBody(text);
|
||||
return true;
|
||||
} catch (error) {
|
||||
return error instanceof Error ? error.message : "Body must be valid JSON";
|
||||
}
|
||||
};
|
||||
|
||||
const buildAuth = (form: KillSwitchFormValue): KillSwitchAuth | null => {
|
||||
switch (form.auth_type) {
|
||||
case "bearer":
|
||||
return { type: "bearer", token: form.auth_token ?? "" };
|
||||
case "api_key":
|
||||
return {
|
||||
type: "api_key",
|
||||
header_name: form.auth_header_name?.trim() || "X-API-Key",
|
||||
api_key: form.auth_api_key ?? "",
|
||||
};
|
||||
case "basic":
|
||||
return { type: "basic", username: form.auth_username ?? "", password: form.auth_password ?? "" };
|
||||
default:
|
||||
return null;
|
||||
}
|
||||
};
|
||||
|
||||
/**
|
||||
* `undefined` means the form never touched the kill switch (leave it as is),
|
||||
* `null` means the user cleared the URL (remove it), otherwise the config to save.
|
||||
*/
|
||||
export const buildKillSwitchFromForm = (form: KillSwitchFormValue | undefined): KillSwitchConfig | null | undefined => {
|
||||
if (form === undefined) return undefined;
|
||||
const url = form.url?.trim() ?? "";
|
||||
if (url.length === 0) return null;
|
||||
return {
|
||||
url,
|
||||
method: form.method ?? "POST",
|
||||
headers: pairsToRecord(form.headers),
|
||||
query_params: pairsToRecord(form.query_params),
|
||||
body: parseKillSwitchBody(form.body),
|
||||
auth: buildAuth(form),
|
||||
};
|
||||
};
|
||||
|
||||
export const parseKillSwitchForForm = (config: KillSwitchConfig | null | undefined): KillSwitchFormValue => {
|
||||
if (!config) return { ...EMPTY_KILL_SWITCH_FORM };
|
||||
const auth = config.auth ?? null;
|
||||
return {
|
||||
url: config.url,
|
||||
method: config.method ?? "POST",
|
||||
headers: recordToPairs(config.headers),
|
||||
query_params: recordToPairs(config.query_params),
|
||||
body: config.body ? JSON.stringify(config.body, null, 2) : "",
|
||||
auth_type: auth?.type ?? "none",
|
||||
auth_token: auth?.type === "bearer" ? auth.token : "",
|
||||
auth_header_name: auth?.type === "api_key" ? auth.header_name : "",
|
||||
auth_api_key: auth?.type === "api_key" ? auth.api_key : "",
|
||||
auth_username: auth?.type === "basic" ? auth.username : "",
|
||||
auth_password: auth?.type === "basic" ? auth.password : "",
|
||||
};
|
||||
};
|
||||
|
|
@ -295,7 +295,7 @@ export const getKeyTableColumns = ({
|
|||
header: () => (
|
||||
<InfoHeader
|
||||
label="Lifetime Spend"
|
||||
tooltip="Cumulative spend across every budget period. Budget resets do not touch this value. Keys created before this field existed only count spend from then on."
|
||||
tooltip="Cumulative spend across every budget period. Budget resets do not touch this value. Lifetime tracking started with LiteLLM v1.103.0 on September 19, 2026, so keys created earlier only count spend since that upgrade."
|
||||
/>
|
||||
),
|
||||
size: 130,
|
||||
|
|
|
|||
|
|
@ -7,6 +7,8 @@ export interface AgentAttachedKey {
|
|||
}
|
||||
|
||||
export type AgentObjectPermission = components["schemas"]["AgentObjectPermission"];
|
||||
export type AgentKillSwitchConfig = components["schemas"]["AgentKillSwitchConfig"];
|
||||
export type AgentKillSwitchResult = components["schemas"]["AgentKillSwitchResult"];
|
||||
|
||||
export interface Agent {
|
||||
agent_id: string;
|
||||
|
|
@ -22,6 +24,7 @@ export interface Agent {
|
|||
};
|
||||
object_permission?: AgentObjectPermission;
|
||||
access_group_ids?: string[] | null;
|
||||
kill_switch?: AgentKillSwitchConfig | null;
|
||||
keys?: AgentAttachedKey[] | null;
|
||||
spend?: number;
|
||||
tpm_limit?: number | null;
|
||||
|
|
|
|||
|
|
@ -6273,6 +6273,16 @@ export const getAgentInfo = async (accessToken: string, agentId: string) => {
|
|||
}
|
||||
};
|
||||
|
||||
export type AgentKillSwitchResult = components["schemas"]["AgentKillSwitchResult"];
|
||||
|
||||
export const triggerAgentKillSwitchCall = async (
|
||||
accessToken: string,
|
||||
agentId: string,
|
||||
): Promise<AgentKillSwitchResult> =>
|
||||
await apiClient.post<AgentKillSwitchResult>(`/v1/agents/${encodeURIComponent(agentId)}/kill_switch`, {
|
||||
accessToken,
|
||||
});
|
||||
|
||||
export const getGuardrailInfo = async (accessToken: string, guardrailId: string) => {
|
||||
try {
|
||||
const url = proxyBaseUrl ? `${proxyBaseUrl}/guardrails/${guardrailId}/info` : `/guardrails/${guardrailId}/info`;
|
||||
|
|
@ -6312,6 +6322,7 @@ export const patchAgentCall = async (
|
|||
session_tpm_limit?: number | null;
|
||||
session_rpm_limit?: number | null;
|
||||
access_group_ids?: string[];
|
||||
kill_switch?: components["schemas"]["AgentKillSwitchConfig"] | null;
|
||||
},
|
||||
) => {
|
||||
try {
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ import { beforeEach, describe, expect, it, vi } from "vitest";
|
|||
import { KeyResponse, Team } from "../key_team_helpers/key_list";
|
||||
import { keyDeleteCall, keyUpdateCall } from "../networking";
|
||||
import { QueryClient } from "@tanstack/react-query";
|
||||
import KeyInfoView from "./key_info_view";
|
||||
import KeyInfoView, { needsLifetimeSpendBackfill } from "./key_info_view";
|
||||
|
||||
const editViewMocks = vi.hoisted(() => ({
|
||||
onSubmit: undefined as ((v: Record<string, any>) => Promise<void>) | undefined,
|
||||
|
|
@ -290,6 +290,58 @@ describe("KeyInfoView", () => {
|
|||
expect(screen.getByTestId("key-lifetime-spend")).toHaveTextContent("Lifetime spend: $340.5000");
|
||||
});
|
||||
|
||||
it("shows the backfill hint when lifetime spend trails the period spend", async () => {
|
||||
vi.mocked(useAuthorized).mockReturnValue(baseUseAuthorizedMock);
|
||||
|
||||
renderWithProviders(
|
||||
<KeyInfoView
|
||||
keyData={{ ...MOCK_KEY_DATA, spend: 10, total_spend: 4 }}
|
||||
onClose={() => {}}
|
||||
keyId={"test-key-id"}
|
||||
onKeyDataUpdate={() => {}}
|
||||
teams={[]}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(await screen.findByTestId("key-lifetime-spend-backfill-hint")).toBeInTheDocument();
|
||||
expect(screen.getByRole("button", { name: /lifetime spend is below/i })).toBeInTheDocument();
|
||||
expect(screen.getByTestId("key-lifetime-spend")).toHaveTextContent("Lifetime spend: $4.0000");
|
||||
});
|
||||
|
||||
it("hides the backfill hint when lifetime spend covers the period spend", async () => {
|
||||
vi.mocked(useAuthorized).mockReturnValue(baseUseAuthorizedMock);
|
||||
|
||||
renderWithProviders(
|
||||
<KeyInfoView
|
||||
keyData={{ ...MOCK_KEY_DATA, spend: 0.25, total_spend: 340.5 }}
|
||||
onClose={() => {}}
|
||||
keyId={"test-key-id"}
|
||||
onKeyDataUpdate={() => {}}
|
||||
teams={[]}
|
||||
/>,
|
||||
);
|
||||
|
||||
expect(await screen.findByTestId("key-lifetime-spend")).toBeInTheDocument();
|
||||
expect(screen.queryByTestId("key-lifetime-spend-backfill-hint")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
describe("needsLifetimeSpendBackfill", () => {
|
||||
it("returns true when total spend is below the period spend", () => {
|
||||
expect(needsLifetimeSpendBackfill(10, 4)).toBe(true);
|
||||
});
|
||||
|
||||
it("returns false when total spend equals or exceeds the period spend", () => {
|
||||
expect(needsLifetimeSpendBackfill(10, 10)).toBe(false);
|
||||
expect(needsLifetimeSpendBackfill(10, 12)).toBe(false);
|
||||
});
|
||||
|
||||
it("treats a missing total spend as zero", () => {
|
||||
expect(needsLifetimeSpendBackfill(10, null)).toBe(true);
|
||||
expect(needsLifetimeSpendBackfill(10, undefined)).toBe(true);
|
||||
expect(needsLifetimeSpendBackfill(0, null)).toBe(false);
|
||||
});
|
||||
});
|
||||
|
||||
it("should render the key's saved router fallbacks", async () => {
|
||||
vi.mocked(useAuthorized).mockReturnValue(baseUseAuthorizedMock);
|
||||
|
||||
|
|
|
|||
|
|
@ -6,11 +6,12 @@ import useTeams from "@/app/(dashboard)/hooks/useTeams";
|
|||
import { useOrganizations } from "@/app/(dashboard)/hooks/organizations/useOrganizations";
|
||||
import { formatNumberWithCommas } from "@/utils/dataUtils";
|
||||
import { mapEmptyStringToNull } from "@/utils/keyUpdateUtils";
|
||||
import { ArrowLeft } from "lucide-react";
|
||||
import { ArrowLeft, Info } from "lucide-react";
|
||||
import { Badge } from "@/components/ui/badge";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { Card } from "@/components/ui/card";
|
||||
import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog";
|
||||
import { HoverCard, HoverCardContent, HoverCardTrigger } from "@/components/ui/hover-card";
|
||||
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
|
||||
import { EntityLink } from "@/components/shared/EntityLink";
|
||||
import { modelGroupHref, teamDetailHref } from "@/utils/entityLinks";
|
||||
|
|
@ -49,6 +50,10 @@ import { parseErrorMessage } from "../shared/errorUtils";
|
|||
import { InheritedBudgetHint, inheritedBudgetGates, keyOwnerBudgetSource } from "../shared/InheritedBudgetHint";
|
||||
import { KeyEditView } from "./key_edit_view";
|
||||
|
||||
export function needsLifetimeSpendBackfill(spend: number, totalSpend: number | null | undefined): boolean {
|
||||
return (totalSpend ?? 0) < spend;
|
||||
}
|
||||
|
||||
interface KeyInfoViewProps {
|
||||
keyId: string;
|
||||
onClose: () => void;
|
||||
|
|
@ -682,6 +687,26 @@ export default function KeyInfoView({
|
|||
)}
|
||||
<p className="text-sm mt-2" data-testid="key-lifetime-spend">
|
||||
Lifetime spend: ${formatNumberWithCommas(currentKeyData.total_spend ?? 0, 4)}
|
||||
{needsLifetimeSpendBackfill(currentKeyData.spend, currentKeyData.total_spend) && (
|
||||
<HoverCard>
|
||||
<HoverCardTrigger
|
||||
render={
|
||||
<button
|
||||
type="button"
|
||||
aria-label="Why lifetime spend is below current spend"
|
||||
className="inline-flex align-middle ml-1 cursor-help"
|
||||
data-testid="key-lifetime-spend-backfill-hint"
|
||||
/>
|
||||
}
|
||||
>
|
||||
<Info className="size-3 text-muted-foreground" />
|
||||
</HoverCardTrigger>
|
||||
<HoverCardContent className="w-80">
|
||||
Lifetime tracking started with LiteLLM v1.103.0 on September 19, 2026 and was not backfilled,
|
||||
so this key's lifetime spend only counts usage since that upgrade.
|
||||
</HoverCardContent>
|
||||
</HoverCard>
|
||||
)}
|
||||
</p>
|
||||
</div>
|
||||
</Card>
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ const ACTION_TONE: Record<string, StatusTone> = {
|
|||
updated: "info",
|
||||
deleted: "error",
|
||||
rotated: "warning",
|
||||
kill_switch_fired: "error",
|
||||
};
|
||||
|
||||
function CopyableJsonBlock({ label, value }: { label: string; value: Record<string, any> }) {
|
||||
|
|
|
|||
|
|
@ -37,6 +37,7 @@ const ACTION_OPTIONS = [
|
|||
{ label: "Updated", value: "updated" },
|
||||
{ label: "Deleted", value: "deleted" },
|
||||
{ label: "Rotated", value: "rotated" },
|
||||
{ label: "Kill switch fired", value: "kill_switch_fired" },
|
||||
] as const;
|
||||
|
||||
const TABLE_OPTIONS = [
|
||||
|
|
@ -45,6 +46,7 @@ const TABLE_OPTIONS = [
|
|||
{ label: "Users", value: "LiteLLM_UserTable" },
|
||||
{ label: "Organizations", value: "LiteLLM_OrganizationTable" },
|
||||
{ label: "Models", value: "LiteLLM_ProxyModelTable" },
|
||||
{ label: "Agents", value: "LiteLLM_AgentsTable" },
|
||||
] as const;
|
||||
|
||||
const ACTION_FILTER_ITEMS = [
|
||||
|
|
|
|||
|
|
@ -24,6 +24,7 @@ export const AUDIT_TABLE_NAME_DISPLAY: Record<string, string> = {
|
|||
LiteLLM_UserTable: "Users",
|
||||
LiteLLM_OrganizationTable: "Organizations",
|
||||
LiteLLM_ProxyModelTable: "Models",
|
||||
LiteLLM_AgentsTable: "Agents",
|
||||
};
|
||||
|
||||
const ACTION_TONE: Record<string, StatusTone> = {
|
||||
|
|
@ -31,9 +32,13 @@ const ACTION_TONE: Record<string, StatusTone> = {
|
|||
updated: "info",
|
||||
deleted: "error",
|
||||
rotated: "warning",
|
||||
kill_switch_fired: "error",
|
||||
};
|
||||
|
||||
const capitalize = (value: string): string => (value ? value.charAt(0).toUpperCase() + value.slice(1) : value);
|
||||
export const auditActionLabel = (action: string): string => {
|
||||
const words = action.replace(/_/g, " ");
|
||||
return words ? words.charAt(0).toUpperCase() + words.slice(1) : words;
|
||||
};
|
||||
|
||||
interface AuditLogsTableColumnsDeps {
|
||||
onViewLog: (log: AuditLogEntry) => void;
|
||||
|
|
@ -55,7 +60,7 @@ export const getAuditLogsTableColumns = ({ onViewLog }: AuditLogsTableColumnsDep
|
|||
size: 110,
|
||||
enableSorting: false,
|
||||
cell: ({ row }) => (
|
||||
<StatusBadge tone={ACTION_TONE[row.original.action] ?? "neutral"} label={capitalize(row.original.action)} />
|
||||
<StatusBadge tone={ACTION_TONE[row.original.action] ?? "neutral"} label={auditActionLabel(row.original.action)} />
|
||||
),
|
||||
},
|
||||
{
|
||||
|
|
|
|||
|
|
@ -59,16 +59,22 @@ export function PluginModeProvider({ children, accessToken }: PluginModeProvider
|
|||
useEffect(() => {
|
||||
// Re-fetch whenever the auth token changes (handles login/logout cycles)
|
||||
if (!accessToken) return;
|
||||
let unmounted = false;
|
||||
pluginApiClient
|
||||
.get("/api/plugins", { accessToken })
|
||||
.then((data: Plugin[]) => {
|
||||
setPlugins(Array.isArray(data) ? data : []);
|
||||
if (!unmounted) setPlugins(Array.isArray(data) ? data : []);
|
||||
})
|
||||
.catch(() => {})
|
||||
// Mark loaded even on failure so a stored plugin mode still falls back to
|
||||
// ai-gateway; otherwise a failed fetch would strand the user on a blank
|
||||
// plugin view with no switcher to escape.
|
||||
.finally(() => setLoaded(true));
|
||||
.finally(() => {
|
||||
if (!unmounted) setLoaded(true);
|
||||
});
|
||||
return () => {
|
||||
unmounted = true;
|
||||
};
|
||||
}, [accessToken]);
|
||||
|
||||
// Once plugins have loaded, fall back to ai-gateway if the persisted mode is
|
||||
|
|
|
|||
283
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
283
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -14615,6 +14615,34 @@ export interface paths {
|
|||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/spend/capture_rate": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
/**
|
||||
* Get Spend Capture Rate
|
||||
* @description Compare the spend LiteLLM captured for a provider against that provider's own bill, per UTC day.
|
||||
*
|
||||
* Admin only. Reads the provider's billing API with the billing credential set on the proxy
|
||||
* (OpenAI: `OPENAI_ADMIN_KEY`) and sums `LiteLLM_DailyUserSpend` for the same days.
|
||||
*
|
||||
* Example:
|
||||
* ```
|
||||
* curl -H "Authorization: Bearer sk-1234" "http://localhost:4000/spend/capture_rate?provider=openai&start_date=2026-09-17&end_date=2026-09-23"
|
||||
* ```
|
||||
*/
|
||||
get: operations["get_spend_capture_rate_spend_capture_rate_get"];
|
||||
put?: never;
|
||||
post?: never;
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/spend/keys": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -18192,6 +18220,37 @@ export interface paths {
|
|||
patch: operations["patch_agent_v1_agents__agent_id__patch"];
|
||||
trace?: never;
|
||||
};
|
||||
"/v1/agents/{agent_id}/kill_switch": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
get?: never;
|
||||
put?: never;
|
||||
/**
|
||||
* Trigger Agent Kill Switch
|
||||
* @description Fire the agent's configured kill switch webhook. Proxy admin only.
|
||||
*
|
||||
* LiteLLM only makes the configured HTTP call and reports what came back; it
|
||||
* does not change the agent's state in LiteLLM. Returns 200 when the webhook
|
||||
* answered 2xx, 502 with the same result body otherwise. Every attempt is
|
||||
* written to the audit log as a `kill_switch_fired` row against the agent.
|
||||
*
|
||||
* Example Request:
|
||||
* ```bash
|
||||
* curl -X POST "http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000/kill_switch" \
|
||||
* -H "Authorization: Bearer <your_api_key>"
|
||||
* ```
|
||||
*/
|
||||
post: operations["trigger_agent_kill_switch_v1_agents__agent_id__kill_switch_post"];
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/v1/agents/{agent_id}/make_public": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -24152,6 +24211,7 @@ export interface components {
|
|||
agent_name: string;
|
||||
/** Extra Headers */
|
||||
extra_headers?: string[] | null;
|
||||
kill_switch?: components["schemas"]["AgentKillSwitchConfig"] | null;
|
||||
/** Litellm Params */
|
||||
litellm_params?: {
|
||||
[key: string]: unknown;
|
||||
|
|
@ -24256,6 +24316,90 @@ export interface components {
|
|||
/** Token */
|
||||
token: string;
|
||||
};
|
||||
/** AgentKillSwitchApiKeyAuth */
|
||||
AgentKillSwitchApiKeyAuth: {
|
||||
/** Api Key */
|
||||
api_key: string;
|
||||
/**
|
||||
* Header Name
|
||||
* @default x-api-key
|
||||
*/
|
||||
header_name: string;
|
||||
/**
|
||||
* @description discriminator enum property added by openapi-typescript
|
||||
* @enum {string}
|
||||
*/
|
||||
type: "api_key";
|
||||
};
|
||||
/** AgentKillSwitchBasicAuth */
|
||||
AgentKillSwitchBasicAuth: {
|
||||
/** Password */
|
||||
password: string;
|
||||
/**
|
||||
* @description discriminator enum property added by openapi-typescript
|
||||
* @enum {string}
|
||||
*/
|
||||
type: "basic";
|
||||
/** Username */
|
||||
username: string;
|
||||
};
|
||||
/** AgentKillSwitchBearerAuth */
|
||||
AgentKillSwitchBearerAuth: {
|
||||
/** Token */
|
||||
token: string;
|
||||
/**
|
||||
* @description discriminator enum property added by openapi-typescript
|
||||
* @enum {string}
|
||||
*/
|
||||
type: "bearer";
|
||||
};
|
||||
/**
|
||||
* AgentKillSwitchConfig
|
||||
* @description Webhook an admin fires to shut an agent down out of band. LiteLLM only
|
||||
* makes the call; whatever the endpoint does with it is the agent's business.
|
||||
*/
|
||||
AgentKillSwitchConfig: {
|
||||
/** Auth */
|
||||
auth?: (components["schemas"]["AgentKillSwitchBearerAuth"] | components["schemas"]["AgentKillSwitchApiKeyAuth"] | components["schemas"]["AgentKillSwitchBasicAuth"]) | null;
|
||||
/** Body */
|
||||
body?: {
|
||||
[key: string]: unknown;
|
||||
} | null;
|
||||
/** Headers */
|
||||
headers?: {
|
||||
[key: string]: string;
|
||||
};
|
||||
/**
|
||||
* Method
|
||||
* @default POST
|
||||
* @enum {string}
|
||||
*/
|
||||
method: "POST" | "PUT" | "PATCH" | "DELETE" | "GET";
|
||||
/** Query Params */
|
||||
query_params?: {
|
||||
[key: string]: string;
|
||||
};
|
||||
/** Url */
|
||||
url: string;
|
||||
};
|
||||
/** AgentKillSwitchResult */
|
||||
AgentKillSwitchResult: {
|
||||
/** Agent Id */
|
||||
agent_id: string;
|
||||
/** Error */
|
||||
error?: string | null;
|
||||
/**
|
||||
* Method
|
||||
* @enum {string}
|
||||
*/
|
||||
method: "POST" | "PUT" | "PATCH" | "DELETE" | "GET";
|
||||
/** Response Body */
|
||||
response_body?: string | null;
|
||||
/** Status Code */
|
||||
status_code?: number | null;
|
||||
/** Url */
|
||||
url: string;
|
||||
};
|
||||
/** AgentMakePublicResponse */
|
||||
AgentMakePublicResponse: {
|
||||
/** Message */
|
||||
|
|
@ -24312,6 +24456,7 @@ export interface components {
|
|||
extra_headers?: string[] | null;
|
||||
/** Keys */
|
||||
keys?: components["schemas"]["AgentKeySummary"][] | null;
|
||||
kill_switch?: components["schemas"]["AgentKillSwitchConfig"] | null;
|
||||
/** Litellm Params */
|
||||
litellm_params?: {
|
||||
[key: string]: unknown;
|
||||
|
|
@ -26742,6 +26887,41 @@ export interface components {
|
|||
*/
|
||||
threshold_step: number;
|
||||
};
|
||||
/** CaptureRateDay */
|
||||
CaptureRateDay: {
|
||||
/** Capture Rate */
|
||||
capture_rate: number | null;
|
||||
/** Captured Spend */
|
||||
captured_spend: number;
|
||||
/** Date */
|
||||
date: string;
|
||||
/** Provider Spend */
|
||||
provider_spend: number;
|
||||
};
|
||||
/** CaptureRateReport */
|
||||
CaptureRateReport: {
|
||||
/** Below Threshold */
|
||||
below_threshold: boolean;
|
||||
/** Capture Rate */
|
||||
capture_rate: number | null;
|
||||
/** Captured Spend */
|
||||
captured_spend: number;
|
||||
/** Days */
|
||||
days: components["schemas"]["CaptureRateDay"][];
|
||||
/** End Date */
|
||||
end_date: string;
|
||||
/**
|
||||
* Provider
|
||||
* @constant
|
||||
*/
|
||||
provider: "openai";
|
||||
/** Provider Spend */
|
||||
provider_spend: number;
|
||||
/** Start Date */
|
||||
start_date: string;
|
||||
/** Threshold */
|
||||
threshold: number;
|
||||
};
|
||||
/** ChangePasswordRequest */
|
||||
ChangePasswordRequest: {
|
||||
/** Current Password */
|
||||
|
|
@ -28273,6 +28453,8 @@ export interface components {
|
|||
reject_clientside_metadata_tags?: boolean | null;
|
||||
/** @description Spreads the proxy's scheduled background jobs (spend flushes, budget resets, config reloads, exports) across a window instead of firing them together on every replica. On by default; set to tune the window, pin a job, or turn it off. */
|
||||
scheduled_job_stagger?: components["schemas"]["ScheduledJobStaggerSettings"] | null;
|
||||
/** @description Daily check of the spend LiteLLM captured against the provider's own bill (OpenAI via OPENAI_ADMIN_KEY). Publishes litellm_spend_capture_rate per provider and alerts when the ratio over the lookback window falls under the threshold (default 0.9). Off unless set. */
|
||||
spend_capture_rate_check?: components["schemas"]["SpendCaptureRateCheckSettings"] | null;
|
||||
/**
|
||||
* Store Model In Db
|
||||
* @description If True, models and config are stored in and loaded from the database. Default is False.
|
||||
|
|
@ -37407,6 +37589,7 @@ export interface components {
|
|||
agent_name?: string;
|
||||
/** Extra Headers */
|
||||
extra_headers?: string[] | null;
|
||||
kill_switch?: components["schemas"]["AgentKillSwitchConfig"] | null;
|
||||
/** Litellm Params */
|
||||
litellm_params?: {
|
||||
[key: string]: unknown;
|
||||
|
|
@ -42252,6 +42435,35 @@ export interface components {
|
|||
/** Model */
|
||||
model?: string | null;
|
||||
};
|
||||
/**
|
||||
* SpendCaptureRateCheckSettings
|
||||
* @description ``general_settings.spend_capture_rate_check``: the daily check of captured spend against the provider bill.
|
||||
*/
|
||||
SpendCaptureRateCheckSettings: {
|
||||
/**
|
||||
* Lookback Days
|
||||
* @default 7
|
||||
*/
|
||||
lookback_days: number;
|
||||
/**
|
||||
* Openai Project Ids
|
||||
* @description Scope the OpenAI bill to these project ids; empty compares against the whole organization. Captured spend is never scoped, so list every project LiteLLM's OpenAI keys belong to
|
||||
* @default []
|
||||
*/
|
||||
openai_project_ids: string[];
|
||||
/**
|
||||
* Providers
|
||||
* @default [
|
||||
* "openai"
|
||||
* ]
|
||||
*/
|
||||
providers: "openai"[];
|
||||
/**
|
||||
* Threshold
|
||||
* @default 0.9
|
||||
*/
|
||||
threshold: number;
|
||||
};
|
||||
/** SpendMetrics */
|
||||
SpendMetrics: {
|
||||
/**
|
||||
|
|
@ -65506,6 +65718,46 @@ export interface operations {
|
|||
};
|
||||
};
|
||||
};
|
||||
get_spend_capture_rate_spend_capture_rate_get: {
|
||||
parameters: {
|
||||
query: {
|
||||
/** @description First UTC day of the range, YYYY-MM-DD */
|
||||
start_date: string;
|
||||
/** @description Last UTC day of the range, YYYY-MM-DD, inclusive */
|
||||
end_date: string;
|
||||
/** @description Provider whose bill to compare against; needs OPENAI_ADMIN_KEY set on the proxy */
|
||||
provider?: "openai";
|
||||
/** @description Ratio under which the report flags below_threshold */
|
||||
threshold?: number;
|
||||
/** @description Scope the OpenAI bill to these project ids; omit to compare against the whole organization. Captured spend is never scoped, so pass every project LiteLLM's OpenAI keys belong to */
|
||||
project_ids?: string[] | null;
|
||||
};
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["CaptureRateReport"];
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
spend_key_fn_spend_keys_get: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -69943,6 +70195,37 @@ export interface operations {
|
|||
};
|
||||
};
|
||||
};
|
||||
trigger_agent_kill_switch_v1_agents__agent_id__kill_switch_post: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path: {
|
||||
agent_id: string;
|
||||
};
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody?: never;
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["AgentKillSwitchResult"];
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
make_agent_public_v1_agents__agent_id__make_public_post: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue