mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
feat(lens): track worker spend through virtual keys (#43989)
* feat(lens): bill worker analysis through virtual keys * fix(lens): pin the verified worker image and add setup proof * fix(lens): preserve network checks and redact billed analysis logs * test(lens): preserve legacy worker result submission during upgrade * fix(lens): enforce trusted worker IPs and restore coverage uploads * docs(lens): explain trusted proxy requirements for worker allowlists * fix(lens): yield to worker disconnects after the synthetic body
This commit is contained in:
parent
6f123b7083
commit
d9f73245be
28 changed files with 1580 additions and 476 deletions
8
.github/workflows/lens-worker.yml
vendored
8
.github/workflows/lens-worker.yml
vendored
|
|
@ -37,7 +37,7 @@ jobs:
|
|||
run: docker build -f deploy/lens/Dockerfile -t lens-worker:${{ github.sha }} .
|
||||
- name: Verify standalone imports with a read-only filesystem
|
||||
run: |
|
||||
docker run --rm --network none --read-only --cap-drop ALL \
|
||||
docker run --rm --network none --read-only --cap-drop ALL --tmpfs /tmp:rw,noexec,nosuid,size=1g \
|
||||
--security-opt no-new-privileges --entrypoint python \
|
||||
lens-worker:${{ github.sha }} -c '
|
||||
import os
|
||||
|
|
@ -47,6 +47,12 @@ jobs:
|
|||
with trace_store() as store:
|
||||
assert store.count() == 0
|
||||
'
|
||||
- name: Verify recovery after temporary storage fills
|
||||
run: |
|
||||
docker run --rm --network none --read-only --cap-drop ALL \
|
||||
--tmpfs /tmp:rw,noexec,nosuid,size=64k --security-opt no-new-privileges \
|
||||
-v "$PWD/tests/proxy_behavior/lens/worker_storage_smoke.py:/app/storage_smoke.py:ro" \
|
||||
--entrypoint python lens-worker:${{ github.sha }} /app/storage_smoke.py
|
||||
- name: Publish versioned Lens worker
|
||||
if: github.event_name != 'pull_request' && github.repository == 'BerriAI/litellm'
|
||||
env:
|
||||
|
|
|
|||
6
.github/workflows/test-postgres.yml
vendored
6
.github/workflows/test-postgres.yml
vendored
|
|
@ -135,7 +135,7 @@ jobs:
|
|||
env:
|
||||
TEST_PATH: ${{ matrix.test-path }}
|
||||
WORKERS: ${{ matrix.workers }}
|
||||
PYTEST_ADDOPTS: ${{ matrix.shard == 'proxy-behavior' && '--cov=litellm/proxy/engine --cov-report=xml:coverage-lens-postgres.xml' || '' }}
|
||||
PYTEST_ADDOPTS: ${{ matrix.shard == 'proxy-behavior' && '--cov=./litellm --cov-report=xml:coverage-lens-postgres.xml' || '' }}
|
||||
run: |
|
||||
if [ "${WORKERS}" = "0" ]; then
|
||||
uv run --no-sync pytest ${TEST_PATH:?} -vv --tb=short --durations=10
|
||||
|
|
@ -145,9 +145,11 @@ jobs:
|
|||
|
||||
- name: Upload Lens database coverage
|
||||
if: steps.changes.outputs.decision != 'skip' && matrix.shard == 'proxy-behavior' && !cancelled()
|
||||
uses: codecov/codecov-action@75cd11691c0faa626561e295848008c8a7dddffe # v5.5.4
|
||||
uses: codecov/codecov-action@303a32d7a59b442fa8d48b6a1cc6825c09c847a5 # v7.1.1
|
||||
with:
|
||||
use_oidc: true
|
||||
version: v11.3.1
|
||||
root_dir: ${{ github.workspace }}
|
||||
files: coverage-lens-postgres.xml
|
||||
flags: lens-postgres
|
||||
fail_ci_if_error: true
|
||||
|
|
|
|||
|
|
@ -2,6 +2,5 @@ FROM python:3.12-slim
|
|||
WORKDIR /app
|
||||
RUN pip install --no-cache-dir httpx==0.28.1 pydantic==2.11.7
|
||||
COPY litellm/proxy/engine/__init__.py litellm/proxy/engine/models.py litellm/proxy/engine/trace_store.py litellm/proxy/engine/analysis.py litellm/proxy/engine/worker.py /app/engine/
|
||||
VOLUME /tmp
|
||||
USER 65532:65532
|
||||
CMD ["python", "-m", "engine.worker"]
|
||||
|
|
|
|||
|
|
@ -6,9 +6,9 @@ Lens reviews recorded activity and saves evidence-linked findings in the LiteLLM
|
|||
|
||||
Upgrade your existing LiteLLM proxy to a release that includes Lens with PostgreSQL, agent tracing (`general_settings.tracing: {store: clickhouse}`), and ClickHouse configured through `CLICKHOUSE_URL` and a separate SELECT-only `CLICKHOUSE_READER_URL`. Enable the ClickHouse callback and request/response logging to analyze LLM requests. Lens can only inspect content you actually retain
|
||||
|
||||
In Lens, click **Connect worker**, then **Generate setup command**. The LiteLLM address is filled in for you; change it only if the server running Docker needs a different network address. Copy the command and run it on your server. The dialog changes to **Worker connected** when the container checks in
|
||||
In Lens, click **Set up analysis**, choose an existing virtual key or **Create worker key**, then **Generate setup command**. The LiteLLM address is filled in for you; change it only if the server running Docker needs a different network address. Copy the command and run it on your server. The dialog changes to **Analyzer connected** when the container checks in
|
||||
|
||||
The command already contains the compatible worker image and one worker token. No separate API key, source checkout, environment file, or second LiteLLM deployment is needed. Keep the command private because it includes the token. The LiteLLM release provides the dashboard and APIs; the container only runs background analysis
|
||||
The command already contains the compatible worker image and one worker token. The selected virtual key stays on the proxy; its secret is never sent to the worker. No source checkout, environment file, or second LiteLLM deployment is needed. Keep the command private because it includes the token. The LiteLLM release provides the dashboard and APIs; the container only runs background analysis
|
||||
|
||||
The dashboard and Compose file pin a verified worker image by digest. The image uses Linux amd64, and the generated command selects that platform. Worker image releases are independent of proxy releases: update the pinned image when changing their API contract. CI also publishes immutable commit tags for reproducible builds
|
||||
|
||||
|
|
@ -20,7 +20,11 @@ docker compose --env-file /path/to/lens.env -f compose.yaml up -d
|
|||
|
||||
Developers can build locally with `LENS_WORKER_IMAGE=litellm-lens-worker:local docker compose -f deploy/lens/compose.yaml -f deploy/lens/compose.build.yaml up -d --build`
|
||||
|
||||
The worker needs outbound HTTPS access to LiteLLM. It needs no inbound ports, provider keys, direct database access, or GPU. The proxy calls your selected model through its configured router; trace content reaches that model provider. Use a model with JSON output support and known token prices. One worker handles one scan at a time and can serve multiple lenses. For more throughput, start another worker with a separate credential
|
||||
The generated command gives the worker 1 GiB of temporary memory-backed storage, shared across parallel reviews. Change `size=1g` in the Docker command or set `LENS_WORKER_TMP_SIZE` with Compose to fit your server and workload. A storage failure marks the scan as failed, cleans up temporary traces, and leaves the worker available for other scans; it does not silently truncate the review. Existing workers must be recreated with the new image and mount options
|
||||
|
||||
The worker needs outbound HTTPS access to LiteLLM. It needs no inbound ports, provider keys, direct database access, or GPU. The proxy calls your selected model through its normal virtual-key authorization and inference pipeline; trace content reaches that model provider. Use a model with JSON output support and known token prices. One worker handles one scan at a time and can serve multiple lenses. For more throughput, start another worker with a separate credential
|
||||
|
||||
If your deployment restricts `allowed_ips`, allow the worker's address. For workers behind a reverse proxy with `use_x_forwarded_for: true`, also configure `mcp_trusted_proxy_ranges` with that proxy's CIDRs and, when needed, `mcp_xff_num_trusted_hops`. Lens reuses these existing trusted-proxy settings. Forwarded addresses without an established trust boundary are rejected by the allowlist; accepting them would let a worker impersonate an allowed address
|
||||
|
||||
V1 setup, manual runs, feedback, and worker credentials are restricted to proxy administrators. Proxy-admin viewers can inspect results. Regular user and team keys cannot access the Lens API. Worker credentials can serve the administrator’s lenses. Revoke it in the connection dialog when retiring a worker. Redeploy the worker alongside proxy upgrades so their API versions match
|
||||
|
||||
|
|
@ -60,7 +64,7 @@ Coverage distinguishes eligible, sampled, reviewed, partial, and unassessable ex
|
|||
|
||||
PostgreSQL stores configurations, findings and all scan history, returned in pages of 50 jobs. Workers claim jobs with optimistic concurrency and a five-minute lease, renewed every 30 seconds. A disconnected job can be reclaimed up to three times. Cancellation stops subsequent work; a model call already in flight may finish and incur cost
|
||||
|
||||
Before every model call, Lens reserves a conservative amount against the monthly lens budget. Successful calls reconcile to reported cost where pricing is available. Interrupted calls retain their reservation because the provider may have charged. A scan stops when the next reservation would exceed the limit, so it can stop with some budget remaining. Lens budgets are separate from virtual-key budgets; analysis calls use the proxy router directly
|
||||
Before every model call, Lens reserves a conservative amount against the monthly lens budget. Successful calls reconcile to reported cost where pricing is available. Interrupted calls retain their reservation because the provider may have charged. A scan stops when the next reservation would exceed the limit, so it can stop with some budget remaining. Both the Lens budget and the selected virtual key’s budgets, model permissions, and rate limits apply. Analysis spend appears under that key in Virtual Keys and normal request logs, with Lens, scan, and worker IDs in request metadata. Analysis prompts and responses are redacted from spend logs; source traces and findings remain available through the administrator-only Lens API. Existing workers need a billing key assigned in **Set up analysis** before they can resume
|
||||
|
||||
V1 requires ClickHouse for both sources. It does not reconstruct sessions from unrelated trace IDs, guarantee exhaustive reviews, cache all per-execution observations across scans, or automatically fix agent code. Trace contents can change as late spans arrive, even though a job's selected IDs are fixed. Findings should be reviewed by a person before acting on them
|
||||
|
||||
|
|
@ -100,6 +104,6 @@ python -m tests.proxy_behavior.lens.evaluate --api-base "$LITELLM_URL" \
|
|||
|
||||
Set `LITELLM_API_KEY` privately. This makes paid model calls. Inspect missed and unexpected per-run labels, final findings and coverage; do not equate a passing dataset with guaranteed detection on arbitrary traces
|
||||
|
||||
The worker uses temporary disk space for trace content while reviewing it, and removes those files after each review. Its Docker image supplies a writable temporary volume while keeping the application filesystem read-only
|
||||
The worker uses temporary disk space for trace content while reviewing it, and removes those files after each review. The Docker command supplies a writable temporary mount while keeping the application filesystem read-only
|
||||
|
||||
To check that accepted behavior stays accepted without hiding new problems, run the evaluator with `--dataset tests/proxy_behavior/lens/feedback_cases.json`. Reports include elapsed time, model call count, reported cost when the proxy provides it, missed checks, unexpected checks, and inconclusive candidates
|
||||
|
|
|
|||
|
|
@ -1,10 +1,12 @@
|
|||
services:
|
||||
lens-worker:
|
||||
image: ${LENS_WORKER_IMAGE:-ghcr.io/berriai/litellm-lens-worker@sha256:40fdb82113dd4474cb6e833cf28552487d87c8baf61693a1c3fc2863b7968c6a}
|
||||
image: ${LENS_WORKER_IMAGE:-ghcr.io/berriai/litellm-lens-worker@sha256:c41e932eaf3e4efbcaf8cc5027c7e93021e5b2823f21cb8785cd107e37b91c9a}
|
||||
environment:
|
||||
LITELLM_URL: ${LITELLM_URL:?Set the URL reachable from this container}
|
||||
LENS_WORKER_TOKEN: ${LENS_WORKER_TOKEN:?Create a worker credential in the Lens UI}
|
||||
restart: unless-stopped
|
||||
read_only: true
|
||||
tmpfs:
|
||||
- /tmp:rw,noexec,nosuid,size=${LENS_WORKER_TMP_SIZE:-1g}
|
||||
cap_drop: [ALL]
|
||||
security_opt: [no-new-privileges:true]
|
||||
|
|
|
|||
BIN
deploy/lens/screenshots/worker-billing-after.png
Normal file
BIN
deploy/lens/screenshots/worker-billing-after.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 59 KiB |
BIN
deploy/lens/screenshots/worker-billing-before.png
Normal file
BIN
deploy/lens/screenshots/worker-billing-before.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 54 KiB |
|
|
@ -33651,6 +33651,18 @@
|
|||
"Worker": {
|
||||
"additionalProperties": false,
|
||||
"properties": {
|
||||
"analysis_key_id": {
|
||||
"anyOf": [
|
||||
{
|
||||
"pattern": "^[a-f0-9]{64}$",
|
||||
"type": "string"
|
||||
},
|
||||
{
|
||||
"type": "null"
|
||||
}
|
||||
],
|
||||
"title": "Analysis Key Id"
|
||||
},
|
||||
"id": {
|
||||
"title": "Id",
|
||||
"type": "string"
|
||||
|
|
@ -33702,6 +33714,11 @@
|
|||
},
|
||||
"WorkerName": {
|
||||
"properties": {
|
||||
"analysis_key_id": {
|
||||
"pattern": "^[a-f0-9]{64}$",
|
||||
"title": "Analysis Key Id",
|
||||
"type": "string"
|
||||
},
|
||||
"name": {
|
||||
"default": "Lens worker",
|
||||
"maxLength": 100,
|
||||
|
|
@ -33710,6 +33727,9 @@
|
|||
"type": "string"
|
||||
}
|
||||
},
|
||||
"required": [
|
||||
"analysis_key_id"
|
||||
],
|
||||
"title": "WorkerName",
|
||||
"type": "object"
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1486,7 +1486,6 @@ async def _user_api_key_auth_builder(
|
|||
general_settings,
|
||||
jwt_handler,
|
||||
litellm_proxy_admin_name,
|
||||
llm_model_list,
|
||||
llm_router,
|
||||
master_key,
|
||||
model_max_budget_limiter,
|
||||
|
|
@ -2181,396 +2180,22 @@ async def _user_api_key_auth_builder(
|
|||
valid_token.end_user_tpd_limit = end_user_params.get("end_user_tpd_limit")
|
||||
valid_token.allowed_model_region = end_user_params.get("allowed_model_region")
|
||||
|
||||
if valid_token is not None:
|
||||
valid_token = _update_key_budget_with_temp_budget_increase(valid_token)
|
||||
|
||||
user_obj: LiteLLM_UserTable | None = None
|
||||
valid_token_dict: dict = {}
|
||||
if valid_token is not None:
|
||||
# Got Valid Token from Cache, DB
|
||||
# Run checks for
|
||||
# 1. If token can call model
|
||||
## 1a. If token can call fallback models (if client-side fallbacks given)
|
||||
# 2. If user_id for this token is in budget
|
||||
# 3. If the user spend within their own team is within budget
|
||||
# 4. If 'user' passed to /chat/completions, /embeddings endpoint is in budget
|
||||
# 5. If token is expired
|
||||
# 6. If token spend is under Budget for the token
|
||||
# 7. If token spend per model is under budget per model
|
||||
# 8. If token spend is under team budget
|
||||
# 9. If team spend is under team budget
|
||||
|
||||
## base case ## key is disabled
|
||||
if valid_token.blocked is True:
|
||||
raise Exception("Key is blocked. Update via `/key/unblock` if you're an admin.")
|
||||
await _enforce_key_and_fallback_model_access(
|
||||
valid_token=valid_token,
|
||||
request_data=request_data,
|
||||
route=route,
|
||||
request=request,
|
||||
llm_model_list=llm_model_list,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
await _prefetch_referenced_auth_objects(
|
||||
valid_token, end_user_id=end_user_id, user_api_key_cache=user_api_key_cache, prisma_client=prisma_client
|
||||
)
|
||||
|
||||
# Check 2. If user_id for this token is in budget - done in common_checks()
|
||||
if valid_token.user_id is not None:
|
||||
try:
|
||||
with tracer.trace("litellm.proxy.auth.get_user_object"):
|
||||
user_obj = await get_user_object(
|
||||
user_id=valid_token.user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
"litellm.proxy.auth.user_api_key_auth.py::user_api_key_auth() - Unable to get user from db/cache. Setting user_obj to None. Exception received - %s",
|
||||
e,
|
||||
)
|
||||
user_obj = None
|
||||
|
||||
if user_obj is not None:
|
||||
# The joint verification-token view carries the key's columns only, so the
|
||||
# user's own per-model budget reaches enforcement and the post-call
|
||||
# increment through the row fetched here.
|
||||
valid_token.user_model_max_budget = user_obj.model_max_budget
|
||||
|
||||
if (
|
||||
user_obj is not None
|
||||
and isinstance(user_obj.metadata, dict)
|
||||
and user_obj.metadata.get("scim_active") is False
|
||||
):
|
||||
raise Exception(
|
||||
f"User={valid_token.user_id} has been deactivated via SCIM. Keys owned by this user cannot be used."
|
||||
)
|
||||
|
||||
# Check 2a. Check if model has zero cost - if so, skip all budget checks
|
||||
model = _get_model_from_request_context(
|
||||
request_data=request_data,
|
||||
route=route,
|
||||
request=request,
|
||||
llm_router=llm_router,
|
||||
team_id=valid_token.team_id,
|
||||
)
|
||||
skip_budget_checks = False
|
||||
if model is not None and llm_router is not None:
|
||||
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
|
||||
|
||||
skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router)
|
||||
if skip_budget_checks:
|
||||
verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model)
|
||||
|
||||
# Check 3. Check if user is in their team budget
|
||||
if not skip_budget_checks and valid_token.team_member_spend is not None:
|
||||
_user_id: Final = valid_token.user_id
|
||||
_team_id: Final = valid_token.team_id
|
||||
if prisma_client is not None and _user_id is not None and _team_id is not None:
|
||||
_cache_key: Final = team_membership_auth_cache_key(team_id=_team_id, user_id=_user_id)
|
||||
|
||||
team_member_info = await user_api_key_cache.async_get_cache(
|
||||
key=_cache_key,
|
||||
model_type=LiteLLM_TeamMembership,
|
||||
)
|
||||
if team_member_info is None:
|
||||
# read from DB
|
||||
_db_member: Final = await TeamMembershipRepository(prisma_client).table.find_first(
|
||||
where={
|
||||
"user_id": _user_id,
|
||||
"team_id": _team_id,
|
||||
},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
if _db_member is not None:
|
||||
team_member_info = LiteLLM_TeamMembership(**_db_member.model_dump())
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=_cache_key,
|
||||
value=team_member_info,
|
||||
model_type=LiteLLM_TeamMembership,
|
||||
ttl=5,
|
||||
)
|
||||
|
||||
if team_member_info is not None and team_member_info.litellm_budget_table is not None:
|
||||
team_member_budget: Final = team_member_info.litellm_budget_table.effective_max_budget(
|
||||
now=datetime.now(timezone.utc),
|
||||
)
|
||||
if team_member_budget is not None and team_member_budget > 0:
|
||||
# Read from cross-pod counter (Redis-first) if available
|
||||
from litellm.proxy.proxy_server import get_current_spend
|
||||
|
||||
team_member_spend = valid_token.team_member_spend
|
||||
if valid_token.user_id is not None and valid_token.team_id is not None:
|
||||
team_member_spend = await get_current_spend(
|
||||
counter_key=f"spend:team_member:{valid_token.user_id}:{valid_token.team_id}",
|
||||
fallback_spend=team_member_spend,
|
||||
max_budget=team_member_budget,
|
||||
)
|
||||
if team_member_spend >= team_member_budget:
|
||||
# common_checks sends this alert on requests that get past here, so only the
|
||||
# request rejected here sends it from the builder.
|
||||
_team_member_max_budget_alert_check(
|
||||
team_id=_team_id,
|
||||
team_alias=valid_token.team_alias,
|
||||
team_metadata=valid_token.team_metadata,
|
||||
organization_id=valid_token.org_id,
|
||||
user_id=_user_id,
|
||||
user_email=user_obj.user_email if user_obj is not None else None,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
spend=team_member_spend,
|
||||
max_budget=team_member_budget,
|
||||
)
|
||||
_entity_id: Final = f"{valid_token.user_id}:{valid_token.team_id}"
|
||||
raise litellm.BudgetExceededError(
|
||||
current_cost=team_member_spend,
|
||||
max_budget=team_member_budget,
|
||||
message=(
|
||||
f"Budget has been exceeded! TeamMember={_entity_id} "
|
||||
f"Current cost: {team_member_spend}, Max budget: {team_member_budget}"
|
||||
),
|
||||
entity_type=Litellm_EntityType.TEAM_MEMBER.value,
|
||||
entity_id=_entity_id,
|
||||
)
|
||||
|
||||
# Check 3. If token is expired
|
||||
if valid_token.expires is not None:
|
||||
current_time = datetime.now(timezone.utc)
|
||||
if isinstance(valid_token.expires, datetime):
|
||||
expiry_time = valid_token.expires
|
||||
else:
|
||||
expiry_time = datetime.fromisoformat(valid_token.expires)
|
||||
if expiry_time.tzinfo is None or expiry_time.tzinfo.utcoffset(expiry_time) is None:
|
||||
expiry_time = expiry_time.replace(tzinfo=timezone.utc)
|
||||
verbose_proxy_logger.debug(
|
||||
"Checking if token expired, expiry time %s and current time %s", expiry_time, current_time
|
||||
)
|
||||
if expiry_time < current_time:
|
||||
# Token exists but is expired.
|
||||
raise ProxyException(
|
||||
message=f"Authentication Error - Expired Key. Key Expiry time {expiry_time} and current time {current_time}",
|
||||
type=ProxyErrorTypes.expired_key,
|
||||
code=status.HTTP_401_UNAUTHORIZED,
|
||||
param=abbreviate_api_key(api_key=api_key),
|
||||
)
|
||||
|
||||
if not skip_budget_checks:
|
||||
with tracer.trace("litellm.proxy.auth.budget_checks"):
|
||||
# Check 4. Max Budget Alert Check (runs before budget enforcement
|
||||
# so multi-threshold 100% alerts fire on the request that crosses
|
||||
# max_budget, before BudgetExceededError is raised below)
|
||||
await _virtual_key_max_budget_alert_check(
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_obj=user_obj,
|
||||
)
|
||||
|
||||
# Check 5. Token Spend is under budget
|
||||
if RouteChecks.is_llm_api_route(route=route):
|
||||
await _virtual_key_max_budget_check(
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_obj=user_obj,
|
||||
)
|
||||
|
||||
# Check 6. Soft Budget Check
|
||||
await _virtual_key_soft_budget_check(
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_obj=user_obj,
|
||||
)
|
||||
|
||||
# Check 5. Token Model Spend is under Model budget
|
||||
max_budget_per_model: Final = valid_token.model_max_budget
|
||||
current_model = _get_model_from_request_context(
|
||||
request_data=request_data,
|
||||
route=route,
|
||||
request=request,
|
||||
llm_router=llm_router,
|
||||
team_id=valid_token.team_id,
|
||||
)
|
||||
current_models = _get_model_names_for_budget_checks(model=current_model)
|
||||
|
||||
if (
|
||||
max_budget_per_model is not None
|
||||
and isinstance(max_budget_per_model, dict)
|
||||
and len(max_budget_per_model) > 0
|
||||
and prisma_client is not None
|
||||
and current_models
|
||||
and valid_token.token is not None
|
||||
):
|
||||
## GET THE SPEND FOR THIS MODEL
|
||||
for model_name in current_models:
|
||||
await _check_key_model_budget_with_fallback(
|
||||
valid_token=valid_token,
|
||||
model_max_budget_limiter=model_max_budget_limiter,
|
||||
model_name=model_name,
|
||||
request_data=request_data,
|
||||
request=request,
|
||||
llm_model_list=llm_model_list,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
# Recompute after a potential budget-fallback rewrite so
|
||||
# the end-user check below validates the final model
|
||||
current_model = _get_model_from_request_context(
|
||||
request_data=request_data,
|
||||
route=route,
|
||||
request=request,
|
||||
llm_router=llm_router,
|
||||
team_id=valid_token.team_id,
|
||||
)
|
||||
current_models = _get_model_names_for_budget_checks(model=current_model)
|
||||
|
||||
# Check 5a. Internal user model_max_budget
|
||||
if current_models:
|
||||
await _check_user_model_budget(
|
||||
valid_token=valid_token,
|
||||
model_max_budget_limiter=model_max_budget_limiter,
|
||||
models=current_models,
|
||||
)
|
||||
|
||||
# Check 5b. End-user model max budget
|
||||
end_user_mmb: Final = valid_token.end_user_model_max_budget
|
||||
if (
|
||||
end_user_mmb is not None
|
||||
and isinstance(end_user_mmb, dict)
|
||||
and len(end_user_mmb) > 0
|
||||
and current_models
|
||||
and valid_token.end_user_id is not None
|
||||
):
|
||||
for model_name in current_models:
|
||||
await model_max_budget_limiter.is_end_user_within_model_budget(
|
||||
end_user_id=valid_token.end_user_id,
|
||||
end_user_model_max_budget=end_user_mmb,
|
||||
model=model_name,
|
||||
)
|
||||
|
||||
# Check 6: Additional Common Checks across jwt + key auth
|
||||
if valid_token.team_id is not None:
|
||||
try:
|
||||
if valid_token.team_id == UI_TEAM_ID:
|
||||
raise TeamNotFoundError(team_id=UI_TEAM_ID)
|
||||
with tracer.trace("litellm.proxy.auth.get_team_object"):
|
||||
_team_obj = await get_team_object(
|
||||
team_id=valid_token.team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except HTTPException:
|
||||
token_team_models: Final = _token_team_models(valid_token)
|
||||
_team_obj = LiteLLM_TeamTableCachedObj(
|
||||
team_id=valid_token.team_id,
|
||||
max_budget=valid_token.team_max_budget,
|
||||
soft_budget=valid_token.team_soft_budget,
|
||||
model_max_budget=valid_token.team_model_max_budget,
|
||||
spend=valid_token.team_spend,
|
||||
tpm_limit=valid_token.team_tpm_limit,
|
||||
rpm_limit=valid_token.team_rpm_limit,
|
||||
tpd_limit=valid_token.team_tpd_limit,
|
||||
blocked=valid_token.team_blocked,
|
||||
models=token_team_models,
|
||||
metadata=valid_token.team_metadata,
|
||||
object_permission_id=valid_token.team_object_permission_id,
|
||||
object_permission=await _resolve_object_permission_for_unresolvable_team(
|
||||
object_permission_id=valid_token.team_object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
),
|
||||
)
|
||||
else:
|
||||
_team_obj = None
|
||||
|
||||
if _team_obj is not None:
|
||||
valid_token.team_object_permission = _team_obj.object_permission
|
||||
# Keep team_metadata in sync with the freshly fetched team so that
|
||||
# guardrails (or any other metadata) added after the key was cached
|
||||
# are picked up on subsequent requests without a cache eviction.
|
||||
valid_token.team_metadata = _team_obj.metadata
|
||||
else:
|
||||
valid_token.team_object_permission = None
|
||||
|
||||
# Fetch project object if key belongs to a project
|
||||
_project_obj = None
|
||||
if valid_token.project_id is not None:
|
||||
_project_obj = await get_project_object(
|
||||
project_id=valid_token.project_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if _project_obj is not None:
|
||||
valid_token.project_metadata = _project_obj.metadata
|
||||
valid_token.project_alias = _project_obj.project_alias
|
||||
|
||||
global_proxy_spend = None
|
||||
if litellm.max_budget > 0 and prisma_client is not None: # user set proxy max budget
|
||||
cache_key: Final = GLOBAL_PROXY_SPEND_CACHE_KEY
|
||||
with tracer.trace("litellm.proxy.auth.get_global_proxy_spend"):
|
||||
global_proxy_spend = await _fetch_global_spend_with_event_coordination(
|
||||
cache_key=cache_key,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
if global_proxy_spend is not None:
|
||||
call_info: Final = CallInfo(
|
||||
token=valid_token.token,
|
||||
spend=global_proxy_spend,
|
||||
max_budget=litellm.max_budget,
|
||||
user_id=litellm_proxy_admin_name,
|
||||
team_id=valid_token.team_id,
|
||||
event_group=Litellm_EntityType.PROXY,
|
||||
)
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.budget_alerts(
|
||||
type="proxy_budget",
|
||||
user_info=call_info,
|
||||
)
|
||||
)
|
||||
# Token passed all checks
|
||||
if valid_token is None:
|
||||
raise HTTPException(401, detail="Invalid API key")
|
||||
if valid_token.token is None:
|
||||
raise HTTPException(401, detail="Invalid API key, no token associated")
|
||||
api_key = valid_token.token
|
||||
|
||||
valid_token_dict = valid_token.model_dump(exclude_none=True)
|
||||
valid_token_dict.pop("token", None)
|
||||
# budget_throttle_pct is excluded from model_dump (it must not leak
|
||||
# into serialized responses), so carry the request-scoped decision
|
||||
# forward by hand to the auth object the rate limiter receives.
|
||||
if valid_token.budget_throttle_pct is not None:
|
||||
valid_token_dict["budget_throttle_pct"] = valid_token.budget_throttle_pct
|
||||
|
||||
if _end_user_object is not None:
|
||||
valid_token_dict.update(end_user_params)
|
||||
valid_token_dict["end_user_object_permission"] = _end_user_object.object_permission
|
||||
|
||||
# check if token is from litellm-ui, litellm ui makes keys to allow users to login with sso. These keys can only be used for LiteLLM UI functions
|
||||
# sso/login, ui/login, /key functions and /user functions
|
||||
# this will never be allowed to call /chat/completions
|
||||
|
||||
if valid_token is None:
|
||||
# No token was found when looking up in the DB
|
||||
raise Exception("Invalid proxy server token passed")
|
||||
if valid_token_dict is not None:
|
||||
virtual_key_auth_obj: Final = await _return_user_api_key_auth_obj(
|
||||
user_obj=user_obj,
|
||||
api_key=api_key,
|
||||
parent_otel_span=parent_otel_span,
|
||||
valid_token_dict=valid_token_dict,
|
||||
route=route,
|
||||
start_time=start_time,
|
||||
)
|
||||
virtual_key_auth_obj.via_virtual_key = True
|
||||
return virtual_key_auth_obj
|
||||
return await validate_resolved_virtual_key(
|
||||
request=request,
|
||||
request_data=cast( # cast-ok: model-alias checks must mutate the original request
|
||||
dict[str, object], request_data
|
||||
),
|
||||
valid_token=valid_token,
|
||||
api_key=api_key,
|
||||
route=route,
|
||||
start_time=start_time,
|
||||
parent_otel_span=parent_otel_span,
|
||||
end_user_id=end_user_id,
|
||||
end_user_params=cast( # cast-ok: builder assembles this dict from validated end-user fields
|
||||
dict[str, object], end_user_params
|
||||
),
|
||||
_end_user_object=_end_user_object,
|
||||
)
|
||||
except Exception as e:
|
||||
return await UserAPIKeyAuthExceptionHandler._handle_authentication_error(
|
||||
e=e,
|
||||
|
|
@ -2583,6 +2208,420 @@ async def _user_api_key_auth_builder(
|
|||
)
|
||||
|
||||
|
||||
async def validate_resolved_virtual_key( # noqa: C901 # Preserve ordering of existing shared authorization checks
|
||||
request: Request,
|
||||
request_data: dict[str, object],
|
||||
valid_token: UserAPIKeyAuth | None,
|
||||
api_key: str,
|
||||
route: str,
|
||||
start_time: datetime,
|
||||
parent_otel_span: Span | None,
|
||||
end_user_id: str | None,
|
||||
end_user_params: dict[str, object],
|
||||
_end_user_object: LiteLLM_EndUserTable | None,
|
||||
) -> UserAPIKeyAuth:
|
||||
from litellm.proxy.proxy_server import (
|
||||
litellm_proxy_admin_name,
|
||||
llm_model_list,
|
||||
llm_router,
|
||||
model_max_budget_limiter,
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
if valid_token is not None:
|
||||
valid_token = _update_key_budget_with_temp_budget_increase(valid_token)
|
||||
|
||||
user_obj: LiteLLM_UserTable | None = None
|
||||
valid_token_dict: dict = {}
|
||||
if valid_token is not None:
|
||||
# Got Valid Token from Cache, DB
|
||||
# Run checks for
|
||||
# 1. If token can call model
|
||||
## 1a. If token can call fallback models (if client-side fallbacks given)
|
||||
# 2. If user_id for this token is in budget
|
||||
# 3. If the user spend within their own team is within budget
|
||||
# 4. If 'user' passed to /chat/completions, /embeddings endpoint is in budget
|
||||
# 5. If token is expired
|
||||
# 6. If token spend is under Budget for the token
|
||||
# 7. If token spend per model is under budget per model
|
||||
# 8. If token spend is under team budget
|
||||
# 9. If team spend is under team budget
|
||||
|
||||
## base case ## key is disabled
|
||||
if valid_token.blocked is True:
|
||||
raise Exception("Key is blocked. Update via `/key/unblock` if you're an admin.")
|
||||
await _enforce_key_and_fallback_model_access(
|
||||
valid_token=valid_token,
|
||||
request_data=request_data,
|
||||
route=route,
|
||||
request=request,
|
||||
llm_model_list=llm_model_list,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
await _prefetch_referenced_auth_objects(
|
||||
valid_token, end_user_id=end_user_id, user_api_key_cache=user_api_key_cache, prisma_client=prisma_client
|
||||
)
|
||||
|
||||
# Check 2. If user_id for this token is in budget - done in common_checks()
|
||||
if valid_token.user_id is not None:
|
||||
try:
|
||||
with tracer.trace("litellm.proxy.auth.get_user_object"):
|
||||
user_obj = await get_user_object(
|
||||
user_id=valid_token.user_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
user_id_upsert=False,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
"litellm.proxy.auth.user_api_key_auth.py::user_api_key_auth() - Unable to get user from db/cache. Setting user_obj to None. Exception received - %s",
|
||||
e,
|
||||
)
|
||||
user_obj = None
|
||||
|
||||
if user_obj is not None:
|
||||
# The joint verification-token view carries the key's columns only, so the
|
||||
# user's own per-model budget reaches enforcement and the post-call
|
||||
# increment through the row fetched here.
|
||||
valid_token.user_model_max_budget = user_obj.model_max_budget
|
||||
|
||||
if (
|
||||
user_obj is not None
|
||||
and isinstance(user_obj.metadata, dict)
|
||||
and user_obj.metadata.get("scim_active") is False
|
||||
):
|
||||
raise Exception(
|
||||
f"User={valid_token.user_id} has been deactivated via SCIM. Keys owned by this user cannot be used."
|
||||
)
|
||||
|
||||
# Check 2a. Check if model has zero cost - if so, skip all budget checks
|
||||
model = _get_model_from_request_context(
|
||||
request_data=request_data,
|
||||
route=route,
|
||||
request=request,
|
||||
llm_router=llm_router,
|
||||
team_id=valid_token.team_id,
|
||||
)
|
||||
skip_budget_checks = False
|
||||
if model is not None and llm_router is not None:
|
||||
from litellm.proxy.auth.auth_checks import _is_model_cost_zero
|
||||
|
||||
skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router)
|
||||
if skip_budget_checks:
|
||||
verbose_proxy_logger.info("Skipping all budget checks for zero-cost model: %s", model)
|
||||
|
||||
# Check 3. Check if user is in their team budget
|
||||
if not skip_budget_checks and valid_token.team_member_spend is not None:
|
||||
_user_id: Final = valid_token.user_id
|
||||
_team_id: Final = valid_token.team_id
|
||||
if prisma_client is not None and _user_id is not None and _team_id is not None:
|
||||
_cache_key: Final = team_membership_auth_cache_key(team_id=_team_id, user_id=_user_id)
|
||||
|
||||
team_member_info = await user_api_key_cache.async_get_cache(
|
||||
key=_cache_key,
|
||||
model_type=LiteLLM_TeamMembership,
|
||||
)
|
||||
if team_member_info is None:
|
||||
# read from DB
|
||||
_db_member: Final = await TeamMembershipRepository(prisma_client).table.find_first(
|
||||
where={
|
||||
"user_id": _user_id,
|
||||
"team_id": _team_id,
|
||||
},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
if _db_member is not None:
|
||||
team_member_info = LiteLLM_TeamMembership(**_db_member.model_dump())
|
||||
await user_api_key_cache.async_set_cache(
|
||||
key=_cache_key,
|
||||
value=team_member_info,
|
||||
model_type=LiteLLM_TeamMembership,
|
||||
ttl=5,
|
||||
)
|
||||
|
||||
if team_member_info is not None and team_member_info.litellm_budget_table is not None:
|
||||
team_member_budget: Final = team_member_info.litellm_budget_table.effective_max_budget(
|
||||
now=datetime.now(timezone.utc),
|
||||
)
|
||||
if team_member_budget is not None and team_member_budget > 0:
|
||||
# Read from cross-pod counter (Redis-first) if available
|
||||
from litellm.proxy.proxy_server import get_current_spend
|
||||
|
||||
team_member_spend = valid_token.team_member_spend
|
||||
if valid_token.user_id is not None and valid_token.team_id is not None:
|
||||
team_member_spend = await get_current_spend(
|
||||
counter_key=f"spend:team_member:{valid_token.user_id}:{valid_token.team_id}",
|
||||
fallback_spend=team_member_spend,
|
||||
max_budget=team_member_budget,
|
||||
)
|
||||
if team_member_spend >= team_member_budget:
|
||||
# common_checks sends this alert on requests that get past here, so only the
|
||||
# request rejected here sends it from the builder.
|
||||
_team_member_max_budget_alert_check(
|
||||
team_id=_team_id,
|
||||
team_alias=valid_token.team_alias,
|
||||
team_metadata=valid_token.team_metadata,
|
||||
organization_id=valid_token.org_id,
|
||||
user_id=_user_id,
|
||||
user_email=user_obj.user_email if user_obj is not None else None,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
spend=team_member_spend,
|
||||
max_budget=team_member_budget,
|
||||
)
|
||||
_entity_id: Final = f"{valid_token.user_id}:{valid_token.team_id}"
|
||||
raise litellm.BudgetExceededError(
|
||||
current_cost=team_member_spend,
|
||||
max_budget=team_member_budget,
|
||||
message=(
|
||||
f"Budget has been exceeded! TeamMember={_entity_id} "
|
||||
f"Current cost: {team_member_spend}, Max budget: {team_member_budget}"
|
||||
),
|
||||
entity_type=Litellm_EntityType.TEAM_MEMBER.value,
|
||||
entity_id=_entity_id,
|
||||
)
|
||||
|
||||
# Check 3. If token is expired
|
||||
if valid_token.expires is not None:
|
||||
current_time = datetime.now(timezone.utc)
|
||||
if isinstance(valid_token.expires, datetime):
|
||||
expiry_time = valid_token.expires
|
||||
else:
|
||||
expiry_time = datetime.fromisoformat(valid_token.expires)
|
||||
if expiry_time.tzinfo is None or expiry_time.tzinfo.utcoffset(expiry_time) is None:
|
||||
expiry_time = expiry_time.replace(tzinfo=timezone.utc)
|
||||
verbose_proxy_logger.debug(
|
||||
"Checking if token expired, expiry time %s and current time %s", expiry_time, current_time
|
||||
)
|
||||
if expiry_time < current_time:
|
||||
# Token exists but is expired.
|
||||
raise ProxyException(
|
||||
message=f"Authentication Error - Expired Key. Key Expiry time {expiry_time} and current time {current_time}",
|
||||
type=ProxyErrorTypes.expired_key,
|
||||
code=status.HTTP_401_UNAUTHORIZED,
|
||||
param=abbreviate_api_key(api_key=api_key),
|
||||
)
|
||||
|
||||
if not skip_budget_checks:
|
||||
with tracer.trace("litellm.proxy.auth.budget_checks"):
|
||||
# Check 4. Max Budget Alert Check (runs before budget enforcement
|
||||
# so multi-threshold 100% alerts fire on the request that crosses
|
||||
# max_budget, before BudgetExceededError is raised below)
|
||||
await _virtual_key_max_budget_alert_check(
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_obj=user_obj,
|
||||
)
|
||||
|
||||
# Check 5. Token Spend is under budget
|
||||
if RouteChecks.is_llm_api_route(route=route):
|
||||
await _virtual_key_max_budget_check(
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_obj=user_obj,
|
||||
)
|
||||
|
||||
# Check 6. Soft Budget Check
|
||||
await _virtual_key_soft_budget_check(
|
||||
valid_token=valid_token,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
user_obj=user_obj,
|
||||
)
|
||||
|
||||
# Check 5. Token Model Spend is under Model budget
|
||||
max_budget_per_model: Final = valid_token.model_max_budget
|
||||
current_model = _get_model_from_request_context(
|
||||
request_data=request_data,
|
||||
route=route,
|
||||
request=request,
|
||||
llm_router=llm_router,
|
||||
team_id=valid_token.team_id,
|
||||
)
|
||||
current_models = _get_model_names_for_budget_checks(model=current_model)
|
||||
|
||||
if (
|
||||
max_budget_per_model is not None
|
||||
and isinstance(max_budget_per_model, dict)
|
||||
and len(max_budget_per_model) > 0
|
||||
and prisma_client is not None
|
||||
and current_models
|
||||
and valid_token.token is not None
|
||||
):
|
||||
## GET THE SPEND FOR THIS MODEL
|
||||
for model_name in current_models:
|
||||
await _check_key_model_budget_with_fallback(
|
||||
valid_token=valid_token,
|
||||
model_max_budget_limiter=model_max_budget_limiter,
|
||||
model_name=model_name,
|
||||
request_data=request_data,
|
||||
request=request,
|
||||
llm_model_list=llm_model_list,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
|
||||
# Recompute after a potential budget-fallback rewrite so
|
||||
# the end-user check below validates the final model
|
||||
current_model = _get_model_from_request_context(
|
||||
request_data=request_data,
|
||||
route=route,
|
||||
request=request,
|
||||
llm_router=llm_router,
|
||||
team_id=valid_token.team_id,
|
||||
)
|
||||
current_models = _get_model_names_for_budget_checks(model=current_model)
|
||||
|
||||
# Check 5a. Internal user model_max_budget
|
||||
if current_models:
|
||||
await _check_user_model_budget(
|
||||
valid_token=valid_token,
|
||||
model_max_budget_limiter=model_max_budget_limiter,
|
||||
models=current_models,
|
||||
)
|
||||
|
||||
# Check 5b. End-user model max budget
|
||||
end_user_mmb: Final = valid_token.end_user_model_max_budget
|
||||
if (
|
||||
end_user_mmb is not None
|
||||
and isinstance(end_user_mmb, dict)
|
||||
and len(end_user_mmb) > 0
|
||||
and current_models
|
||||
and valid_token.end_user_id is not None
|
||||
):
|
||||
for model_name in current_models:
|
||||
await model_max_budget_limiter.is_end_user_within_model_budget(
|
||||
end_user_id=valid_token.end_user_id,
|
||||
end_user_model_max_budget=end_user_mmb,
|
||||
model=model_name,
|
||||
)
|
||||
|
||||
# Check 6: Additional Common Checks across jwt + key auth
|
||||
if valid_token.team_id is not None:
|
||||
try:
|
||||
if valid_token.team_id == UI_TEAM_ID:
|
||||
raise TeamNotFoundError(team_id=UI_TEAM_ID)
|
||||
with tracer.trace("litellm.proxy.auth.get_team_object"):
|
||||
_team_obj = await get_team_object(
|
||||
team_id=valid_token.team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except HTTPException:
|
||||
token_team_models: Final = _token_team_models(valid_token)
|
||||
_team_obj = LiteLLM_TeamTableCachedObj(
|
||||
team_id=valid_token.team_id,
|
||||
max_budget=valid_token.team_max_budget,
|
||||
soft_budget=valid_token.team_soft_budget,
|
||||
model_max_budget=valid_token.team_model_max_budget,
|
||||
spend=valid_token.team_spend,
|
||||
tpm_limit=valid_token.team_tpm_limit,
|
||||
rpm_limit=valid_token.team_rpm_limit,
|
||||
tpd_limit=valid_token.team_tpd_limit,
|
||||
blocked=valid_token.team_blocked,
|
||||
models=token_team_models,
|
||||
metadata=valid_token.team_metadata,
|
||||
object_permission_id=valid_token.team_object_permission_id,
|
||||
object_permission=await _resolve_object_permission_for_unresolvable_team(
|
||||
object_permission_id=valid_token.team_object_permission_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
),
|
||||
)
|
||||
else:
|
||||
_team_obj = None
|
||||
|
||||
if _team_obj is not None:
|
||||
valid_token.team_object_permission = _team_obj.object_permission
|
||||
# Keep team_metadata in sync with the freshly fetched team so that
|
||||
# guardrails (or any other metadata) added after the key was cached
|
||||
# are picked up on subsequent requests without a cache eviction.
|
||||
valid_token.team_metadata = _team_obj.metadata
|
||||
else:
|
||||
valid_token.team_object_permission = None
|
||||
|
||||
# Fetch project object if key belongs to a project
|
||||
_project_obj = None
|
||||
if valid_token.project_id is not None:
|
||||
_project_obj = await get_project_object(
|
||||
project_id=valid_token.project_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if _project_obj is not None:
|
||||
valid_token.project_metadata = _project_obj.metadata
|
||||
valid_token.project_alias = _project_obj.project_alias
|
||||
|
||||
global_proxy_spend = None
|
||||
if litellm.max_budget > 0 and prisma_client is not None: # user set proxy max budget
|
||||
cache_key: Final = GLOBAL_PROXY_SPEND_CACHE_KEY
|
||||
with tracer.trace("litellm.proxy.auth.get_global_proxy_spend"):
|
||||
global_proxy_spend = await _fetch_global_spend_with_event_coordination(
|
||||
cache_key=cache_key,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
if global_proxy_spend is not None:
|
||||
call_info: Final = CallInfo(
|
||||
token=valid_token.token,
|
||||
spend=global_proxy_spend,
|
||||
max_budget=litellm.max_budget,
|
||||
user_id=litellm_proxy_admin_name,
|
||||
team_id=valid_token.team_id,
|
||||
event_group=Litellm_EntityType.PROXY,
|
||||
)
|
||||
asyncio.create_task(
|
||||
proxy_logging_obj.budget_alerts(
|
||||
type="proxy_budget",
|
||||
user_info=call_info,
|
||||
)
|
||||
)
|
||||
# Token passed all checks
|
||||
if valid_token is None:
|
||||
raise HTTPException(401, detail="Invalid API key")
|
||||
if valid_token.token is None:
|
||||
raise HTTPException(401, detail="Invalid API key, no token associated")
|
||||
api_key = valid_token.token
|
||||
|
||||
valid_token_dict = valid_token.model_dump(exclude_none=True)
|
||||
valid_token_dict.pop("token", None)
|
||||
# budget_throttle_pct is excluded from model_dump (it must not leak
|
||||
# into serialized responses), so carry the request-scoped decision
|
||||
# forward by hand to the auth object the rate limiter receives.
|
||||
if valid_token.budget_throttle_pct is not None:
|
||||
valid_token_dict["budget_throttle_pct"] = valid_token.budget_throttle_pct
|
||||
|
||||
if _end_user_object is not None:
|
||||
valid_token_dict.update(end_user_params)
|
||||
valid_token_dict["end_user_object_permission"] = _end_user_object.object_permission
|
||||
|
||||
# check if token is from litellm-ui, litellm ui makes keys to allow users to login with sso. These keys can only be used for LiteLLM UI functions
|
||||
# sso/login, ui/login, /key functions and /user functions
|
||||
# this will never be allowed to call /chat/completions
|
||||
|
||||
if valid_token is None:
|
||||
# No token was found when looking up in the DB
|
||||
raise Exception("Invalid proxy server token passed")
|
||||
if valid_token_dict is not None:
|
||||
virtual_key_auth_obj: Final = await _return_user_api_key_auth_obj(
|
||||
user_obj=user_obj,
|
||||
api_key=api_key,
|
||||
parent_otel_span=parent_otel_span,
|
||||
valid_token_dict=valid_token_dict,
|
||||
route=route,
|
||||
start_time=start_time,
|
||||
)
|
||||
virtual_key_auth_obj.via_virtual_key = True
|
||||
return virtual_key_auth_obj
|
||||
|
||||
|
||||
async def _safe_fetch(label: str, awaitable):
|
||||
"""Run an awaitable and return its result. Re-raises authentication /
|
||||
authorization failures (HTTPException, ProxyException,
|
||||
|
|
@ -2718,6 +2757,8 @@ async def _run_centralized_common_checks(
|
|||
request: Request,
|
||||
request_data: dict[str, object],
|
||||
route: str,
|
||||
*,
|
||||
force_virtual_key_checks: bool = False,
|
||||
) -> None:
|
||||
"""Run ``common_checks`` once at the ``user_api_key_auth`` wrapper
|
||||
boundary, regardless of which ``_user_api_key_auth_builder`` path
|
||||
|
|
@ -2751,7 +2792,9 @@ async def _run_centralized_common_checks(
|
|||
# auth in the builder — the wrapper must not retroactively apply
|
||||
# authz on top, or k8s readiness probes and other unauthenticated
|
||||
# callers get 401.
|
||||
if route in LiteLLMRoutes.public_routes.value or route_in_additonal_public_routes(current_route=route):
|
||||
if not force_virtual_key_checks and (
|
||||
route in LiteLLMRoutes.public_routes.value or route_in_additonal_public_routes(current_route=route)
|
||||
):
|
||||
return
|
||||
|
||||
# User-configured pass-through endpoints with ``auth: false`` are
|
||||
|
|
@ -2761,7 +2804,7 @@ async def _run_centralized_common_checks(
|
|||
# admin-only. The "auth" flag on the endpoint config is the
|
||||
# contract; honor it.
|
||||
pass_through_endpoints: Final = general_settings.get("pass_through_endpoints", None)
|
||||
if pass_through_endpoints is not None:
|
||||
if not force_virtual_key_checks and pass_through_endpoints is not None:
|
||||
for endpoint in pass_through_endpoints:
|
||||
if isinstance(endpoint, dict) and endpoint.get("path", "") == route and endpoint.get("auth") is not True:
|
||||
return
|
||||
|
|
@ -2772,10 +2815,14 @@ async def _run_centralized_common_checks(
|
|||
# Running common_checks would block every admin route on these
|
||||
# deployments where that was previously not the contract. If any
|
||||
# authn is enabled (JWT, OAuth2, OAuth2-proxy), authz must run.
|
||||
if is_no_auth_dev_mode(master_key, general_settings):
|
||||
if not force_virtual_key_checks and is_no_auth_dev_mode(master_key, general_settings):
|
||||
return
|
||||
|
||||
if user_custom_auth is not None and not general_settings.get("custom_auth_run_common_checks", False):
|
||||
if (
|
||||
not force_virtual_key_checks
|
||||
and user_custom_auth is not None
|
||||
and not general_settings.get("custom_auth_run_common_checks", False)
|
||||
):
|
||||
return
|
||||
|
||||
parent_otel_span: Final = user_api_key_auth_obj.parent_otel_span
|
||||
|
|
@ -3204,6 +3251,8 @@ async def _authorize_authenticated_request(
|
|||
request_data: dict,
|
||||
route: str,
|
||||
api_key: str,
|
||||
*,
|
||||
force_virtual_key_checks: bool = False,
|
||||
) -> UserAPIKeyAuth | None:
|
||||
"""Authorize an already-authenticated request: disabled-route check, the single
|
||||
``common_checks`` gate (which also reserves budget), and end-user fallback
|
||||
|
|
@ -3264,6 +3313,7 @@ async def _authorize_authenticated_request(
|
|||
request=request,
|
||||
request_data=authorized_data,
|
||||
route=route,
|
||||
force_virtual_key_checks=force_virtual_key_checks,
|
||||
)
|
||||
except Exception as e:
|
||||
return await UserAPIKeyAuthExceptionHandler._handle_authentication_error(
|
||||
|
|
@ -3883,3 +3933,45 @@ async def _run_post_custom_auth_checks(
|
|||
valid_token.project_alias = _project_obj.project_alias
|
||||
|
||||
return valid_token
|
||||
|
||||
|
||||
async def authorize_internal_virtual_key(
|
||||
key_hash: str, request: Request, request_data: dict[str, object]
|
||||
) -> UserAPIKeyAuth:
|
||||
"""Authorize a server-owned job against its persisted virtual-key assignment, never a client-supplied bearer hash."""
|
||||
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
|
||||
|
||||
identity: Final = IdentityStore.key_from_principal(
|
||||
await IdentityStore(prisma_client, user_api_key_cache, proxy_logging_obj=proxy_logging_obj).resolve(
|
||||
hashed_token=key_hash
|
||||
)
|
||||
)
|
||||
route: Final = get_request_route(request=request)
|
||||
await pre_db_read_auth_checks(request_data=request_data, request=request, route=route)
|
||||
auth: Final = await validate_resolved_virtual_key(
|
||||
request=request,
|
||||
request_data=request_data,
|
||||
valid_token=identity,
|
||||
api_key=key_hash,
|
||||
route=route,
|
||||
start_time=datetime.now(timezone.utc),
|
||||
parent_otel_span=None,
|
||||
end_user_id=None,
|
||||
end_user_params={}, # mutable-ok: existing end-user validation contract
|
||||
_end_user_object=None,
|
||||
)
|
||||
auth.budget_reservation = None
|
||||
recovered: Final = await _authorize_authenticated_request(
|
||||
user_api_key_auth_obj=auth,
|
||||
request=request,
|
||||
request_data=request_data,
|
||||
route=route,
|
||||
api_key=key_hash,
|
||||
force_virtual_key_checks=True,
|
||||
)
|
||||
if recovered is not None:
|
||||
return recovered
|
||||
_seed_request_destinations(auth, request)
|
||||
auth.request_route = route
|
||||
request.state.principal = _resolve_request_principal(request, auth)
|
||||
return auth
|
||||
|
|
|
|||
98
litellm/proxy/engine/billing.py
Normal file
98
litellm/proxy/engine/billing.py
Normal file
|
|
@ -0,0 +1,98 @@
|
|||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from typing import Final
|
||||
|
||||
import orjson
|
||||
from fastapi import HTTPException, Request, Response
|
||||
from pydantic import TypeAdapter
|
||||
from starlette.types import Message
|
||||
|
||||
import litellm
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
from litellm.proxy.auth.resolvers.store import IdentityStore
|
||||
from litellm.proxy.auth.user_api_key_auth import authorize_internal_virtual_key
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.proxy.spend_tracking.budget_reservation import release_unbound_budget_reservation
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
||||
async def validate_key(key_id: str | None) -> UserAPIKeyAuth | None:
|
||||
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
|
||||
|
||||
if key_id is None:
|
||||
return None
|
||||
key: Final = IdentityStore.key_from_principal(
|
||||
await IdentityStore(prisma_client, user_api_key_cache, proxy_logging_obj=proxy_logging_obj).resolve(
|
||||
hashed_token=key_id
|
||||
)
|
||||
)
|
||||
if key.blocked or key.is_session_token:
|
||||
raise HTTPException(400, "Choose an active virtual key for Lens analysis")
|
||||
return key
|
||||
|
||||
|
||||
async def complete(
|
||||
key_id: str, data: dict[str, object], reserve: Callable[[], Awaitable[None]], incoming: Request
|
||||
) -> tuple[ModelResponse, float | None]:
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy.proxy_server import llm_router, proxy_config, proxy_logging_obj, version
|
||||
|
||||
payload: Final = orjson.dumps(data)
|
||||
client_ip: Final = IPAddressUtils.get_mcp_client_ip(incoming)
|
||||
|
||||
body: Final[Message] = {
|
||||
"type": "http.request",
|
||||
"body": payload,
|
||||
"more_body": False,
|
||||
}
|
||||
messages: Final = iter((body,))
|
||||
|
||||
async def receive() -> Message:
|
||||
message: Final = next(messages, None)
|
||||
return message if message is not None else await incoming.receive()
|
||||
|
||||
request: Final = Request(
|
||||
{ # mutable-ok: Starlette mutates its ASGI scope
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/v1/chat/completions",
|
||||
"raw_path": b"/v1/chat/completions",
|
||||
"query_string": b"",
|
||||
"headers": [(b"content-type", b"application/json")], # mutable-ok: ASGI header contract
|
||||
"scheme": incoming.url.scheme or "http",
|
||||
"client": (client_ip, incoming.client.port if incoming.client else 0) if client_ip else None,
|
||||
"server": ("litellm.internal", 80),
|
||||
},
|
||||
receive=receive,
|
||||
)
|
||||
try:
|
||||
auth: Final = await authorize_internal_virtual_key(key_id, request, data)
|
||||
await reserve()
|
||||
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
|
||||
fastapi_response: Final = Response()
|
||||
try:
|
||||
response: Final = TypeAdapter(ModelResponse).validate_python(
|
||||
await processor.base_process_llm_request(
|
||||
request=request,
|
||||
fastapi_response=fastapi_response,
|
||||
user_api_key_dict=auth,
|
||||
route_type="acompletion",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
general_settings=TypeAdapter(dict[str, object]).validate_python(proxy_server.general_settings), # pyright: ignore[reportUnknownMemberType] # Validate the legacy untyped config at the request boundary
|
||||
proxy_config=proxy_config,
|
||||
llm_router=llm_router,
|
||||
version=version,
|
||||
)
|
||||
)
|
||||
billed: Final = fastapi_response.headers.get("x-litellm-response-cost")
|
||||
return response, float(billed) if billed not in (None, "", "None") else None
|
||||
except Exception as exc:
|
||||
raise await processor._handle_llm_api_exception( # pyright: ignore[reportPrivateUsage] # Standard proxy endpoint failure hook releases limits and records failures
|
||||
e=exc, user_api_key_dict=auth, proxy_logging_obj=proxy_logging_obj, version=version
|
||||
)
|
||||
except litellm.BudgetExceededError:
|
||||
raise HTTPException(402, "The analysis key or its owner has reached a budget limit")
|
||||
finally:
|
||||
reservation: Final = getattr(request.state, "budget_reservation", None)
|
||||
if isinstance(reservation, Mapping):
|
||||
await release_unbound_budget_reservation(TypeAdapter(dict[str, object]).validate_python(reservation))
|
||||
|
|
@ -6,13 +6,14 @@ from types import MappingProxyType
|
|||
from typing import Annotated, Final, TypeAlias
|
||||
from uuid import uuid4
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
||||
from pydantic import AwareDatetime, BaseModel, Field, TypeAdapter
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.db.routing_prisma_wrapper import writer_wrapper
|
||||
from litellm.proxy.engine.billing import validate_key
|
||||
from litellm.proxy.engine.models import (
|
||||
Claim,
|
||||
Engine,
|
||||
|
|
@ -277,28 +278,52 @@ async def preview_sample(body: Preview, auth: Auth) -> Sample:
|
|||
)
|
||||
|
||||
|
||||
class WorkerName(BaseModel):
|
||||
class WorkerBilling(BaseModel):
|
||||
analysis_key_id: str = Field(pattern=r"^[a-f0-9]{64}$")
|
||||
|
||||
|
||||
class WorkerName(WorkerBilling):
|
||||
name: str = Field(default="Lens worker", min_length=1, max_length=100)
|
||||
|
||||
|
||||
@router.post("/workers/register", response_model=WorkerCreated)
|
||||
async def register_worker(body: WorkerName, auth: Auth) -> WorkerCreated:
|
||||
scope: Final = user_scope(auth, write=True)
|
||||
await validate_key(body.analysis_key_id)
|
||||
token: Final = "lens-" + secrets.token_urlsafe(40)
|
||||
worker: Final = Worker(
|
||||
id=str(uuid4()), name=body.name, scope=scope, last_seen=datetime(1970, 1, 1, tzinfo=timezone.utc)
|
||||
id=str(uuid4()),
|
||||
name=body.name,
|
||||
scope=scope,
|
||||
analysis_key_id=body.analysis_key_id,
|
||||
last_seen=datetime(1970, 1, 1, tzinfo=timezone.utc),
|
||||
)
|
||||
await repository().save_worker(worker, hashlib.sha256(token.encode()).hexdigest())
|
||||
return WorkerCreated(worker=worker, token=token)
|
||||
|
||||
|
||||
@router.put("/workers/{worker_id}/billing-key", response_model=Worker)
|
||||
async def set_worker_billing(worker_id: str, body: WorkerBilling, auth: Auth) -> Worker:
|
||||
scope: Final = user_scope(auth, write=True)
|
||||
worker: Final = next((w for w in await repository().workers() if w.id == worker_id), None)
|
||||
if worker is None or not can_access(scope, worker.scope):
|
||||
raise HTTPException(404, "Worker not found")
|
||||
if worker.revoked:
|
||||
raise HTTPException(409, "Register a new worker instead of updating revoked access")
|
||||
await validate_key(body.analysis_key_id)
|
||||
updated: Final = await repository().set_worker_billing(worker.id, body.analysis_key_id)
|
||||
if updated is None:
|
||||
raise HTTPException(409, "Worker access was revoked")
|
||||
return updated
|
||||
|
||||
|
||||
@router.delete("/workers/{worker_id}")
|
||||
async def revoke_worker(worker_id: str, auth: Auth) -> bool:
|
||||
scope: Final = user_scope(auth, write=True)
|
||||
worker: Final = next((w for w in await repository().workers() if w.id == worker_id), None)
|
||||
if worker is None or not can_access(scope, worker.scope):
|
||||
raise HTTPException(404, "Worker not found")
|
||||
await repository().save_worker(worker.model_copy(update=MappingProxyType({"revoked": True})))
|
||||
await repository().revoke_worker(worker.id)
|
||||
return True
|
||||
|
||||
|
||||
|
|
@ -306,6 +331,8 @@ async def revoke_worker(worker_id: str, auth: Auth) -> bool:
|
|||
async def claim(worker: WorkerAuth, protocol_version: int = 1) -> Claim | None:
|
||||
if protocol_version != 2:
|
||||
raise HTTPException(409, "Upgrade the Lens worker using the current Connect worker command")
|
||||
if worker.analysis_key_id is None:
|
||||
raise HTTPException(409, "Assign an analysis key to this worker in Lens setup")
|
||||
now: Final = datetime.now(timezone.utc)
|
||||
await repository().heartbeat(worker.id, now.isoformat())
|
||||
for candidate in await repository().engines():
|
||||
|
|
@ -398,11 +425,11 @@ async def content(
|
|||
|
||||
|
||||
@router.post("/worker/{engine_id}/{job_id}/model", response_model=ModelResult)
|
||||
async def model(engine_id: str, job_id: str, body: ModelRequest, worker: WorkerAuth) -> ModelResult:
|
||||
async def model(engine_id: str, job_id: str, body: ModelRequest, worker: WorkerAuth, request: Request) -> ModelResult:
|
||||
from litellm.proxy.engine.inference import analyze
|
||||
|
||||
engine, job = await assigned(engine_id, job_id, worker)
|
||||
return await analyze(repository(), engine, job, worker.id, body)
|
||||
return await analyze(repository(), engine, job, worker, body, request)
|
||||
|
||||
|
||||
@router.post("/worker/{engine_id}/{job_id}/result", response_model=Engine)
|
||||
|
|
@ -487,7 +514,9 @@ async def claim_candidate(candidate: Engine, worker: Worker, now: datetime) -> C
|
|||
scheduled: Final = queue_job(e, now, job_id) if e.settings.enabled and e.next_run_at <= now else e
|
||||
return claim_job(scheduled, worker, now)
|
||||
|
||||
updated: Final = required(await repository().update(candidate.id, schedule))
|
||||
updated: Final = await repository().update(candidate.id, schedule, changed_only=True)
|
||||
if updated is None:
|
||||
return None
|
||||
job: Final = current_job(updated)
|
||||
if job and job.worker_id == worker.id and job.status == "running" and job != current_job(candidate):
|
||||
return Claim(engine_id=updated.id, job=job, findings=updated.findings)
|
||||
|
|
|
|||
|
|
@ -2,12 +2,14 @@ from datetime import datetime, timezone
|
|||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from fastapi import HTTPException
|
||||
from fastapi import HTTPException, Request
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.clickhouse.context import lens_analysis
|
||||
from litellm.proxy.engine.models import Engine, Job, ModelRequest, ModelResult
|
||||
from litellm.litellm_core_utils.initialize_dynamic_callback_params import inherit_message_logging_privacy
|
||||
from litellm.proxy.engine.billing import complete, validate_key
|
||||
from litellm.proxy.engine.models import Engine, Job, ModelRequest, ModelResult, Worker
|
||||
from litellm.proxy.engine.repository import EngineRepository
|
||||
from litellm.proxy.engine.state import current_job, renew_budget, replace_job
|
||||
from litellm.types.utils import CostPerToken, ModelResponse
|
||||
|
|
@ -84,14 +86,20 @@ def quote(deployments: tuple[Deployment, ...], prompt: str) -> float:
|
|||
return ((len((prompt + _SYSTEM).encode()) + 1024) * input_rate + 4096 * output_rate) * 2
|
||||
|
||||
|
||||
async def analyze(repo: EngineRepository, engine: Engine, job: Job, worker_id: str, body: ModelRequest) -> ModelResult:
|
||||
async def analyze(
|
||||
repo: EngineRepository, engine: Engine, job: Job, worker: Worker, body: ModelRequest, request: Request
|
||||
) -> ModelResult:
|
||||
from litellm.proxy.proxy_server import llm_router
|
||||
|
||||
if llm_router is None:
|
||||
raise HTTPException(503, "No analysis models are configured")
|
||||
if worker.analysis_key_id is None:
|
||||
raise HTTPException(409, "Assign an analysis key to this worker in Lens setup")
|
||||
billing_key: Final = await validate_key(worker.analysis_key_id)
|
||||
team_id: Final = billing_key.team_id if billing_key else None
|
||||
deployments: Final = tuple(
|
||||
Deployment.model_validate(d)
|
||||
for d in llm_router.get_model_list(model_name=job.settings.model, team_id=engine.scope.team_id or None) or ()
|
||||
for d in llm_router.get_model_list(model_name=job.settings.model, team_id=team_id) or ()
|
||||
)
|
||||
if not deployments:
|
||||
raise HTTPException(400, "Analysis model is no longer available")
|
||||
|
|
@ -101,7 +109,13 @@ async def analyze(repo: EngineRepository, engine: Engine, job: Job, worker_id: s
|
|||
def reserve(e: Engine) -> Engine:
|
||||
current: Final = renew_budget(e, now)
|
||||
active: Final = current_job(current)
|
||||
if active is None or active.id != job.id or active.worker_id != worker_id:
|
||||
if (
|
||||
active is None
|
||||
or active.id != job.id
|
||||
or active.worker_id != worker.id
|
||||
or active.lease_until is None
|
||||
or active.lease_until <= datetime.now(timezone.utc)
|
||||
):
|
||||
raise HTTPException(409, "Job was cancelled or reassigned")
|
||||
if current.spent + estimate > current.settings.monthly_budget:
|
||||
raise HTTPException(402, "Monthly lens budget reached; increase it or wait for next month")
|
||||
|
|
@ -109,28 +123,35 @@ async def analyze(repo: EngineRepository, engine: Engine, job: Job, worker_id: s
|
|||
current, active.model_copy(update=MappingProxyType({"cost": active.cost + estimate}))
|
||||
).model_copy(update=MappingProxyType({"spent": current.spent + estimate}))
|
||||
|
||||
if await repo.update(engine.id, reserve) is None:
|
||||
raise HTTPException(409, "Could not reserve analysis budget")
|
||||
with lens_analysis():
|
||||
response: Final = await llm_router.acompletion( # pyright: ignore[reportUnknownMemberType] # Router forwards provider-specific keyword arguments
|
||||
model=job.settings.model,
|
||||
messages=[ # mutable-ok: Router requires OpenAI message dictionaries in a list
|
||||
{"role": "system", "content": _SYSTEM}, # mutable-ok: provider message dictionary
|
||||
{"role": "user", "content": body.prompt}, # mutable-ok: provider message dictionary
|
||||
],
|
||||
max_tokens=4096,
|
||||
stream=False,
|
||||
timeout=120,
|
||||
num_retries=0,
|
||||
disable_fallbacks=True,
|
||||
response_format={"type": "json_object"}, # mutable-ok: provider response-format JSON object
|
||||
metadata={ # mutable-ok: Router mutates metadata
|
||||
"tags": ["litellm-engine"], # mutable-ok: logging callbacks require a tag list
|
||||
"user_api_key_team_id": engine.scope.team_id,
|
||||
},
|
||||
)
|
||||
async def reserve_budget() -> None:
|
||||
if await repo.update(engine.id, reserve) is None:
|
||||
raise HTTPException(409, "Could not reserve analysis budget")
|
||||
|
||||
data: Final[dict[str, object]] = { # mutable-ok: proxy processing enriches request data
|
||||
"model": job.settings.model,
|
||||
"messages": [ # mutable-ok: OpenAI request contract
|
||||
{"role": "system", "content": _SYSTEM}, # mutable-ok: OpenAI message contract
|
||||
{"role": "user", "content": body.prompt}, # mutable-ok: OpenAI message contract
|
||||
],
|
||||
"max_tokens": 4096,
|
||||
"stream": False,
|
||||
"timeout": 120,
|
||||
"num_retries": 0,
|
||||
"disable_fallbacks": True,
|
||||
"response_format": {"type": "json_object"}, # mutable-ok: provider response-format JSON
|
||||
"metadata": { # mutable-ok: request processing enriches metadata
|
||||
"tags": ["litellm-engine"], # mutable-ok: logging callbacks require a list
|
||||
"lens_id": engine.id,
|
||||
"lens_run_id": job.id,
|
||||
"lens_worker_id": worker.id,
|
||||
"user_api_key_team_id": team_id,
|
||||
},
|
||||
}
|
||||
|
||||
with lens_analysis(), inherit_message_logging_privacy(True):
|
||||
response, billed_cost = await complete(worker.analysis_key_id, data, reserve_budget, request)
|
||||
parsed: Final = Completion.model_validate_json(response.model_dump_json())
|
||||
cost: Final = completion_charge(deployments, response, estimate)
|
||||
cost: Final = billed_cost if billed_cost is not None else completion_charge(deployments, response, estimate)
|
||||
|
||||
def settle(e: Engine) -> Engine:
|
||||
charged: Final = next((j for j in e.jobs if j.id == job.id), None)
|
||||
|
|
|
|||
|
|
@ -193,6 +193,7 @@ class Engine(Record):
|
|||
|
||||
|
||||
class Worker(Record):
|
||||
analysis_key_id: str | None = Field(default=None, pattern=r"^[a-f0-9]{64}$")
|
||||
id: str
|
||||
name: str
|
||||
scope: Scope
|
||||
|
|
|
|||
|
|
@ -45,20 +45,24 @@ class EngineRepository:
|
|||
)
|
||||
return engine
|
||||
|
||||
async def update(self, engine_id: str, transform: Callable[[Engine], Engine], attempts: int = 8) -> Engine | None:
|
||||
async def update(
|
||||
self, engine_id: str, transform: Callable[[Engine], Engine], attempts: int = 8, *, changed_only: bool = False
|
||||
) -> Engine | None:
|
||||
for _ in range(attempts):
|
||||
completed, updated = await self._try_update(engine_id, transform)
|
||||
completed, updated = await self._try_update(engine_id, transform, changed_only)
|
||||
if completed:
|
||||
return updated
|
||||
return None
|
||||
|
||||
async def _try_update(self, engine_id: str, transform: Callable[[Engine], Engine]) -> tuple[bool, Engine | None]:
|
||||
async def _try_update(
|
||||
self, engine_id: str, transform: Callable[[Engine], Engine], changed_only: bool
|
||||
) -> tuple[bool, Engine | None]:
|
||||
previous: Final = await self.get(engine_id)
|
||||
if previous is None:
|
||||
return True, None
|
||||
candidate: Final = transform(previous)
|
||||
if candidate == previous:
|
||||
return True, previous
|
||||
return True, None if changed_only else previous
|
||||
updated: Final = candidate.model_copy(update=MappingProxyType({"version": previous.version + 1}))
|
||||
rows: Final = _ROWS.validate_python(
|
||||
await self.db.query_raw(
|
||||
|
|
@ -135,6 +139,24 @@ class EngineRepository:
|
|||
'UPDATE "LiteLLM_EngineWorker" SET data=$1::jsonb WHERE id=$2', worker.model_dump_json(), worker.id
|
||||
)
|
||||
|
||||
async def set_worker_billing(self, worker_id: str, key_id: str) -> Worker | None:
|
||||
rows: Final = _ROWS.validate_python(
|
||||
await self.db.query_raw(
|
||||
"""UPDATE "LiteLLM_EngineWorker"
|
||||
SET data=jsonb_set(data, '{analysis_key_id}', to_jsonb($1::text))
|
||||
WHERE id=$2 AND COALESCE((data->>'revoked')::boolean, false)=false RETURNING data""",
|
||||
key_id,
|
||||
worker_id,
|
||||
)
|
||||
)
|
||||
return Worker.model_validate(rows[0].data) if rows else None
|
||||
|
||||
async def revoke_worker(self, worker_id: str) -> None:
|
||||
await self.db.execute_raw(
|
||||
"""UPDATE "LiteLLM_EngineWorker" SET data=jsonb_set(data, '{revoked}', 'true') WHERE id=$1""",
|
||||
worker_id,
|
||||
)
|
||||
|
||||
async def heartbeat(self, worker_id: str, now: str) -> None:
|
||||
await self.db.execute_raw(
|
||||
"""UPDATE "LiteLLM_EngineWorker" SET data=jsonb_set(data, '{last_seen}', to_jsonb($1::text)) WHERE id=$2""",
|
||||
|
|
|
|||
|
|
@ -91,8 +91,17 @@ def renew_budget(engine: Engine, now: datetime) -> Engine:
|
|||
|
||||
|
||||
def merge_finding(engine: Engine, draft: FindingDraft, revision: int, now: datetime) -> Finding:
|
||||
identity: Final = hashlib.sha256(f"{engine.id}:{draft.check_id}:{draft.title.lower()}".encode()).hexdigest()[:24]
|
||||
previous: Final = next((f for f in engine.findings if f.id == (draft.existing_finding_id or identity)), None)
|
||||
legacy_identity: Final = hashlib.sha256(f"{engine.id}:{draft.check_id}:{draft.title.lower()}".encode()).hexdigest()[
|
||||
:24
|
||||
]
|
||||
identity: Final = hashlib.sha256(
|
||||
f"{engine.id}:{draft.check_id}:{draft.kind}:{draft.title.lower()}".encode()
|
||||
).hexdigest()[:24]
|
||||
identities: Final = (draft.existing_finding_id, identity, legacy_identity)
|
||||
previous: Final = next(
|
||||
(f for f in engine.findings if f.id in identities and f.kind == draft.kind and f.check_id == draft.check_id),
|
||||
None,
|
||||
)
|
||||
occurrences: Final = tuple(sorted(frozenset(e.execution_id for e in draft.evidence if e.role == "support")))
|
||||
if previous is None:
|
||||
return Finding(
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import sqlite3
|
||||
from collections.abc import Awaitable, Callable
|
||||
from contextlib import suppress
|
||||
from types import MappingProxyType
|
||||
|
|
@ -82,10 +83,12 @@ class EngineWorker:
|
|||
result: Final = await analyze_sample(claim, sample, read, model, progress)
|
||||
saved: Final = await self.client.post(prefix + "/result", json=result.model_dump(mode="json"))
|
||||
saved.raise_for_status()
|
||||
except (httpx.HTTPError, ValueError) as exc:
|
||||
except (httpx.HTTPError, ValueError, OSError, sqlite3.Error) as exc:
|
||||
status: Final = exc.response.status_code if isinstance(exc, httpx.HTTPStatusError) else None
|
||||
message: Final = (
|
||||
"Monthly budget reached"
|
||||
"Worker temporary storage failed. Increase its capacity or reduce analysis parallelism."
|
||||
if isinstance(exc, (OSError, sqlite3.Error))
|
||||
else "Monthly budget reached"
|
||||
if status == 402
|
||||
else "Analysis interrupted. Check worker connectivity, model configuration, and trace storage."
|
||||
)
|
||||
|
|
|
|||
203
tests/integration/spend/test_lens_billing.py
Normal file
203
tests/integration/spend/test_lens_billing.py
Normal file
|
|
@ -0,0 +1,203 @@
|
|||
import json
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from hashlib import sha256
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
from tests.integration._support.client import Gateway, eventually, object_value, string_value
|
||||
from tests.integration._support.database import read_rows, write_rows
|
||||
from tests.integration._support.process import owned_proxy
|
||||
from tests.integration.pricing.test_off_peak_pricing import off_peak_window
|
||||
|
||||
|
||||
def delete_lens(engine_id: str) -> None:
|
||||
write_rows('DELETE FROM "LiteLLM_EngineRun" WHERE engine_id=%s', (engine_id,))
|
||||
write_rows('DELETE FROM "LiteLLM_Engine" WHERE id=%s', (engine_id,))
|
||||
assert read_rows('SELECT id FROM "LiteLLM_Engine" WHERE id=%s', (engine_id,)) == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize("off_peak", (False, True))
|
||||
def test_lens_bills_selected_key_and_rechecks_its_permissions(gateway: Gateway, off_peak: bool) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
input_cost_per_token=0.000001,
|
||||
output_cost_per_token=0.000002,
|
||||
model_info={
|
||||
"off_peak_pricing": {
|
||||
**off_peak_window(-1, 1),
|
||||
"input_cost_per_token": 0.0000005,
|
||||
"output_cost_per_token": 0.000001,
|
||||
}
|
||||
}
|
||||
if off_peak
|
||||
else None,
|
||||
)
|
||||
key: Final = scenario.key(models=[model], max_budget=1)
|
||||
key_id: Final = sha256(key.encode()).hexdigest()
|
||||
worker: Final = gateway.post(
|
||||
"/engine/workers/register", {"name": "Billing regression", "analysis_key_id": key_id}
|
||||
)
|
||||
worker_id: Final = string_value(object_value(worker["worker"])["id"])
|
||||
scenario.cleanups.callback(write_rows, 'DELETE FROM "LiteLLM_EngineWorker" WHERE id=%s', (worker_id,))
|
||||
engine: Final = gateway.post(
|
||||
"/engine",
|
||||
{
|
||||
"name": "Billing regression",
|
||||
"model": model,
|
||||
"enabled": False,
|
||||
"context": "Answers should be accurate",
|
||||
"source": "requests",
|
||||
},
|
||||
)
|
||||
engine_id: Final = string_value(engine["id"])
|
||||
scenario.cleanups.callback(delete_lens, engine_id)
|
||||
worker_key: Final = string_value(worker["token"])
|
||||
unauthorized: Final = gateway.request(
|
||||
"POST", "/engine/workers/register", {"name": "Denied", "analysis_key_id": key_id}, key=key
|
||||
)
|
||||
assert unauthorized.status_code == 403, unauthorized.text
|
||||
with ThreadPoolExecutor(max_workers=8) as pool:
|
||||
claims: Final = tuple(
|
||||
pool.map(
|
||||
lambda _: gateway.request("POST", "/engine/worker/claim?protocol_version=2", {}, key=worker_key),
|
||||
range(8),
|
||||
)
|
||||
)
|
||||
assert all(response.status_code == 200 for response in claims)
|
||||
winners: Final = tuple(response.json() for response in claims if response.json() is not None)
|
||||
assert len(winners) == 1
|
||||
claim: Final = object_value(winners[0])
|
||||
assert claim["engine_id"] == engine_id
|
||||
job_id: Final = string_value(object_value(claim["job"])["id"])
|
||||
path: Final = f"/engine/worker/{engine_id}/{job_id}/model"
|
||||
result: Final = gateway.post(path, {"prompt": "Inspect this run", "purpose": "extract"}, key=worker_key)
|
||||
expected: Final = (20 * 0.000001 + 20 * 0.000002) * (0.5 if off_peak else 1)
|
||||
assert result["cost"] == pytest.approx(expected)
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows('SELECT spend FROM "LiteLLM_VerificationToken" WHERE token=%s', (key_id,)),
|
||||
lambda values: len(values) == 1 and values[0]["spend"] == pytest.approx(expected),
|
||||
seconds=70,
|
||||
)
|
||||
assert rows[0]["spend"] == pytest.approx(expected)
|
||||
assert gateway.get(f"/engine/{engine_id}")["spent"] == pytest.approx(expected)
|
||||
raw_hash: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": "Not a bearer credential"}],
|
||||
},
|
||||
key=key_id,
|
||||
)
|
||||
assert raw_hash.status_code == 401, raw_hash.text
|
||||
gateway.post("/key/update", {"key": key, "max_budget": expected / 2})
|
||||
exhausted: Final = gateway.request(
|
||||
"POST", path, {"prompt": "Must not run", "purpose": "extract"}, key=worker_key
|
||||
)
|
||||
assert exhausted.status_code == 402, exhausted.text
|
||||
gateway.post("/key/update", {"key": key, "max_budget": 1, "models": ["unavailable-analysis-model"]})
|
||||
restricted: Final = gateway.request(
|
||||
"POST", path, {"prompt": "Must not run", "purpose": "extract"}, key=worker_key
|
||||
)
|
||||
assert restricted.status_code == 403, restricted.text
|
||||
gateway.post("/key/block", {"key": key})
|
||||
blocked: Final = gateway.request("POST", path, {"prompt": "Must not run", "purpose": "extract"}, key=worker_key)
|
||||
assert blocked.status_code == 400, blocked.text
|
||||
assert gateway.get(f"/engine/{engine_id}")["spent"] == pytest.approx(expected)
|
||||
replacement: Final = scenario.key(models=[model], rpm_limit=1)
|
||||
replacement_id: Final = sha256(replacement.encode()).hexdigest()
|
||||
changed: Final = gateway.request(
|
||||
"PUT", f"/engine/workers/{worker_id}/billing-key", {"analysis_key_id": replacement_id}
|
||||
)
|
||||
assert changed.status_code == 200, changed.text
|
||||
billed_replacement: Final = gateway.post(
|
||||
path, {"prompt": "Inspect another run", "purpose": "extract"}, key=worker_key
|
||||
)
|
||||
assert billed_replacement["cost"] == pytest.approx(expected)
|
||||
limited: Final = gateway.request("POST", path, {"prompt": "Must not run", "purpose": "extract"}, key=worker_key)
|
||||
assert limited.status_code == 429, limited.text
|
||||
second_rows: Final = eventually(
|
||||
lambda: read_rows('SELECT spend FROM "LiteLLM_VerificationToken" WHERE token=%s', (replacement_id,)),
|
||||
lambda values: len(values) == 1 and values[0]["spend"] == pytest.approx(expected),
|
||||
seconds=70,
|
||||
)
|
||||
assert second_rows[0]["spend"] == pytest.approx(expected)
|
||||
revoked: Final = gateway.request("DELETE", f"/engine/workers/{worker_id}")
|
||||
assert revoked.status_code == 200, revoked.text
|
||||
denied_worker: Final = gateway.request(
|
||||
"POST", path, {"prompt": "Must not run", "purpose": "extract"}, key=worker_key
|
||||
)
|
||||
assert denied_worker.status_code == 401, denied_worker.text
|
||||
forbidden_change: Final = gateway.request(
|
||||
"PUT", f"/engine/workers/{worker_id}/billing-key", {"analysis_key_id": replacement_id}
|
||||
)
|
||||
assert forbidden_change.status_code == 409, forbidden_change.text
|
||||
gateway.post(f"/engine/{engine_id}/cancel", {})
|
||||
|
||||
|
||||
@pytest.mark.parametrize("cancel_on_disconnect", (False, True))
|
||||
def test_worker_spend_logs_do_not_expose_investigation_content(
|
||||
gateway: Gateway, tmp_path: Path, cancel_on_disconnect: bool
|
||||
) -> None:
|
||||
config: Final = tmp_path / "lens-privacy.json"
|
||||
config.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"model_list": [],
|
||||
"general_settings": {
|
||||
"master_key": "os.environ/LITELLM_MASTER_KEY",
|
||||
"database_url": "os.environ/DATABASE_URL",
|
||||
"store_model_in_db": True,
|
||||
"store_prompts_in_spend_logs": True,
|
||||
"cancel_on_disconnect": cancel_on_disconnect,
|
||||
"proxy_batch_write_at": 1,
|
||||
"proxy_batch_polling_interval": 1,
|
||||
"allowed_ips": ["127.0.0.1"],
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
with owned_proxy(gateway, tmp_path, {}, config=config) as isolated, isolated.scenario() as scenario:
|
||||
model: Final = scenario.model(input_cost_per_token=0.000001, output_cost_per_token=0.000002)
|
||||
key: Final = scenario.key(models=[model])
|
||||
key_id: Final = sha256(key.encode()).hexdigest()
|
||||
marker: Final = "PRIVATE_OTHER_TEAM_TRACE_CONTENT"
|
||||
ordinary: Final = isolated.chat(model, key=key, text=marker)
|
||||
retained: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT proxy_server_request FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
|
||||
(string_value(ordinary["id"]),),
|
||||
),
|
||||
lambda rows: len(rows) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
assert marker in str(retained[0]), "Control must prove this proxy retains ordinary prompts"
|
||||
worker: Final = isolated.post("/engine/workers/register", {"analysis_key_id": key_id})
|
||||
worker_id: Final = string_value(object_value(worker["worker"])["id"])
|
||||
scenario.cleanups.callback(write_rows, 'DELETE FROM "LiteLLM_EngineWorker" WHERE id=%s', (worker_id,))
|
||||
engine: Final = isolated.post(
|
||||
"/engine", {"name": "Log privacy", "model": model, "enabled": False, "context": "Find problems"}
|
||||
)
|
||||
engine_id: Final = string_value(engine["id"])
|
||||
scenario.cleanups.callback(delete_lens, engine_id)
|
||||
worker_token: Final = string_value(worker["token"])
|
||||
claim: Final = isolated.post("/engine/worker/claim?protocol_version=2", {}, key=worker_token)
|
||||
job_id: Final = string_value(object_value(claim["job"])["id"])
|
||||
result: Final = isolated.post(
|
||||
f"/engine/worker/{engine_id}/{job_id}/model", {"prompt": marker, "purpose": "extract"}, key=worker_token
|
||||
)
|
||||
assert result["content"], "The worker must still receive model output"
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT spend, proxy_server_request, response FROM "LiteLLM_SpendLogs" WHERE api_key=%s AND request_id<>%s',
|
||||
(key_id, string_value(ordinary["id"])),
|
||||
),
|
||||
lambda rows: len(rows) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
assert float(rows[0]["spend"]) == pytest.approx(result["cost"])
|
||||
assert marker not in str(rows[0])
|
||||
assert result["content"] not in str(rows[0]["response"])
|
||||
isolated.post(f"/engine/{engine_id}/cancel", {})
|
||||
|
|
@ -1,12 +1,14 @@
|
|||
import asyncio
|
||||
import hashlib
|
||||
import os
|
||||
from collections.abc import AsyncIterator
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Final
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from fastapi import HTTPException
|
||||
from fastapi import HTTPException, Request
|
||||
from fastapi.security import HTTPAuthorizationCredentials
|
||||
|
||||
from litellm import Router
|
||||
|
|
@ -22,6 +24,14 @@ from litellm.proxy.utils import PrismaClient, ProxyLogging
|
|||
async def lens_database() -> AsyncIterator[PrismaClient]:
|
||||
original_db: Final = proxy_server.prisma_client
|
||||
original_router: Final = proxy_server.llm_router
|
||||
original_settings: Final = proxy_server.general_settings
|
||||
proxy_server.general_settings = {
|
||||
**original_settings,
|
||||
"allowed_ips": ["127.0.0.1"],
|
||||
"use_x_forwarded_for": True,
|
||||
"mcp_trusted_proxy_ranges": ["192.0.2.100/32"],
|
||||
"mcp_xff_num_trusted_hops": 1,
|
||||
}
|
||||
client: Final = PrismaClient(os.environ["DATABASE_URL"], ProxyLogging(UserApiKeyCache()))
|
||||
await client.connect()
|
||||
proxy_server.prisma_client = client
|
||||
|
|
@ -42,6 +52,7 @@ async def lens_database() -> AsyncIterator[PrismaClient]:
|
|||
try:
|
||||
yield client
|
||||
finally:
|
||||
proxy_server.general_settings = original_settings
|
||||
proxy_server.prisma_client = original_db
|
||||
proxy_server.llm_router = original_router
|
||||
await client.disconnect()
|
||||
|
|
@ -57,7 +68,11 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database:
|
|||
checks=(Check(id="retries", instruction="Find unrecovered retries"),),
|
||||
)
|
||||
engine: Final = await endpoints.create_engine(settings, admin)
|
||||
registration: Final = await endpoints.register_worker(endpoints.WorkerName(name="Test analyzer"), admin)
|
||||
key_id: Final = hashlib.sha256(uuid4().bytes).hexdigest()
|
||||
await lens_database.db.litellm_verificationtoken.create(data={"token": key_id, "models": ["lens-test-analysis"]})
|
||||
registration: Final = await endpoints.register_worker(
|
||||
endpoints.WorkerName(name="Test analyzer", analysis_key_id=key_id), admin
|
||||
)
|
||||
credentials: Final = HTTPAuthorizationCredentials(scheme="Bearer", credentials=registration.token)
|
||||
worker: Final = await endpoints.worker_auth(credentials)
|
||||
try:
|
||||
|
|
@ -70,8 +85,12 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database:
|
|||
listing: Final = await endpoints.list_engines(admin)
|
||||
assert engine.id in tuple(e.id for e in listing.engines)
|
||||
assert worker.id in tuple(w.id for w in listing.workers)
|
||||
claimed: Final = await endpoints.claim_candidate(engine, worker, datetime.now(timezone.utc))
|
||||
assert claimed is not None
|
||||
claims: Final = await asyncio.gather(
|
||||
*(endpoints.claim_candidate(engine, worker, datetime.now(timezone.utc)) for _ in range(8))
|
||||
)
|
||||
winners: Final = tuple(claim for claim in claims if claim is not None)
|
||||
assert len(winners) == 1
|
||||
claimed: Final = winners[0]
|
||||
assert claimed.job.worker_id == worker.id
|
||||
assert (
|
||||
await endpoints.claim_candidate(
|
||||
|
|
@ -88,13 +107,80 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database:
|
|||
claimed.job.id,
|
||||
ModelRequest(prompt="Return an empty observations list", purpose="extract"),
|
||||
worker,
|
||||
Request(
|
||||
{
|
||||
"type": "http",
|
||||
"scheme": "http",
|
||||
"path": "/engine/worker/model",
|
||||
"headers": [],
|
||||
"client": ("127.0.0.1", 1234),
|
||||
}
|
||||
),
|
||||
)
|
||||
assert '"observations"' in response.content
|
||||
with pytest.raises(HTTPException) as denied_ip:
|
||||
await endpoints.model(
|
||||
engine.id,
|
||||
claimed.job.id,
|
||||
ModelRequest(prompt="Must not run", purpose="extract"),
|
||||
worker,
|
||||
Request(
|
||||
{
|
||||
"type": "http",
|
||||
"scheme": "http",
|
||||
"path": "/engine/worker/model",
|
||||
"headers": [(b"x-forwarded-for", b"127.0.0.1")],
|
||||
"client": ("192.0.2.1", 1234),
|
||||
}
|
||||
),
|
||||
)
|
||||
assert denied_ip.value.status_code == 403
|
||||
forwarded: Final = await endpoints.model(
|
||||
engine.id,
|
||||
claimed.job.id,
|
||||
ModelRequest(prompt="Return an empty observations list", purpose="extract"),
|
||||
worker,
|
||||
Request(
|
||||
{
|
||||
"type": "http",
|
||||
"scheme": "http",
|
||||
"path": "/engine/worker/model",
|
||||
"headers": [(b"x-forwarded-for", b"127.0.0.1")],
|
||||
"client": ("192.0.2.100", 1234),
|
||||
}
|
||||
),
|
||||
)
|
||||
assert '"observations"' in forwarded.content
|
||||
with pytest.raises(HTTPException) as spoofed_chain:
|
||||
await endpoints.model(
|
||||
engine.id,
|
||||
claimed.job.id,
|
||||
ModelRequest(prompt="Must not run", purpose="extract"),
|
||||
worker,
|
||||
Request(
|
||||
{
|
||||
"type": "http",
|
||||
"scheme": "http",
|
||||
"path": "/engine/worker/model",
|
||||
"headers": [(b"x-forwarded-for", b"127.0.0.1, 192.0.2.1")],
|
||||
"client": ("192.0.2.100", 1234),
|
||||
}
|
||||
),
|
||||
)
|
||||
assert spoofed_chain.value.status_code == 403
|
||||
charged: Final = await endpoints.get_engine(engine.id, worker.scope)
|
||||
assert charged.spent == pytest.approx(response.cost)
|
||||
assert charged.jobs[0].cost == pytest.approx(response.cost)
|
||||
assert charged.spent == pytest.approx(response.cost + forwarded.cost)
|
||||
assert charged.jobs[0].cost == pytest.approx(response.cost + forwarded.cost)
|
||||
legacy: Final = worker.model_copy(update={"analysis_key_id": None})
|
||||
await endpoints.repository().save_worker(legacy)
|
||||
authenticated_legacy: Final = await endpoints.worker_auth(credentials)
|
||||
assert authenticated_legacy.analysis_key_id is None
|
||||
with pytest.raises(HTTPException) as needs_billing:
|
||||
await endpoints.claim(authenticated_legacy, protocol_version=2)
|
||||
assert needs_billing.value.status_code == 409
|
||||
assert await endpoints.heartbeat(engine.id, claimed.job.id, authenticated_legacy)
|
||||
finished: Final = await endpoints.result(
|
||||
engine.id, claimed.job.id, Result(coverage=Coverage(screened=2)), worker
|
||||
engine.id, claimed.job.id, Result(coverage=Coverage(screened=2)), authenticated_legacy
|
||||
)
|
||||
assert finished.jobs[0].status == "completed"
|
||||
assert finished.jobs[0].coverage.screened == 2
|
||||
|
|
@ -124,6 +210,10 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database:
|
|||
assert cancelled.jobs[0].status == "cancelled"
|
||||
assert await endpoints.cancel_engine(engine.id, admin) == cancelled
|
||||
assert await endpoints.revoke_worker(worker.id, admin)
|
||||
assert await endpoints.repository().set_worker_billing(worker.id, key_id) is None
|
||||
with pytest.raises(HTTPException) as revoked_billing:
|
||||
await endpoints.set_worker_billing(worker.id, endpoints.WorkerBilling(analysis_key_id=key_id), admin)
|
||||
assert revoked_billing.value.status_code == 409
|
||||
with pytest.raises(HTTPException) as revoked:
|
||||
await endpoints.worker_auth(credentials)
|
||||
assert revoked.value.status_code == 401
|
||||
|
|
@ -134,3 +224,4 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database:
|
|||
await lens_database.db.execute_raw('DELETE FROM "LiteLLM_EngineRun" WHERE engine_id=$1', engine.id)
|
||||
await lens_database.db.execute_raw('DELETE FROM "LiteLLM_Engine" WHERE id=$1', engine.id)
|
||||
await lens_database.db.execute_raw('DELETE FROM "LiteLLM_EngineWorker" WHERE id=$1', worker.id)
|
||||
await lens_database.db.execute_raw('DELETE FROM "LiteLLM_VerificationToken" WHERE token=$1', key_id)
|
||||
|
|
|
|||
88
tests/proxy_behavior/lens/worker_storage_smoke.py
Normal file
88
tests/proxy_behavior/lens/worker_storage_smoke.py
Normal file
|
|
@ -0,0 +1,88 @@
|
|||
import asyncio
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from engine.models import (
|
||||
Claim,
|
||||
EngineSettings,
|
||||
Execution,
|
||||
ExecutionContent,
|
||||
Job,
|
||||
ModelResult,
|
||||
Result,
|
||||
Sample,
|
||||
TracePart,
|
||||
)
|
||||
from engine.worker import EngineWorker
|
||||
|
||||
|
||||
async def main() -> None:
|
||||
now: Final = datetime(2026, 1, 1, tzinfo=timezone.utc)
|
||||
claims: Final = iter(("full", "healthy"))
|
||||
saved: Final = SimpleQueue[Result]()
|
||||
pages: Final = SimpleQueue[str]()
|
||||
settings: Final = EngineSettings(name="Storage recovery", model="unused", context="Finish the task", concurrency=1)
|
||||
execution: Final = Execution(
|
||||
id="run", source="traces", trace_id="trace", team_id="", name="Task", start_time="", span_count=10000
|
||||
)
|
||||
|
||||
def handle(request: httpx.Request) -> httpx.Response:
|
||||
path: Final = request.url.path
|
||||
if path.endswith("/claim"):
|
||||
claim: Final = Claim(
|
||||
engine_id="lens",
|
||||
job=Job(id=next(claims), created_at=now, start=now, end=now, settings=settings, revision=1),
|
||||
findings=(),
|
||||
)
|
||||
return httpx.Response(200, json=claim.model_dump(mode="json"))
|
||||
if path.endswith("/sample"):
|
||||
return httpx.Response(200, json=Sample(executions=(execution,), eligible=1).model_dump())
|
||||
if path.endswith("/content"):
|
||||
healthy: Final = "/healthy/" in path
|
||||
cursor: Final = request.url.params.get("cursor", "")
|
||||
pages.put(cursor)
|
||||
assert pages.qsize() < 100, "The deliberately small temporary mount must fill"
|
||||
content: Final = ExecutionContent(
|
||||
execution=execution,
|
||||
parts=tuple(
|
||||
TracePart(
|
||||
execution_id="run",
|
||||
span_id=f"{cursor}-{i}",
|
||||
name="tool",
|
||||
kind="tool",
|
||||
content="Finished" if healthy else "x" * 8000,
|
||||
)
|
||||
for i in range(1 if healthy else 40)
|
||||
),
|
||||
next_cursor=None if healthy else str(pages.qsize()),
|
||||
)
|
||||
return httpx.Response(200, json=content.model_dump())
|
||||
if path.endswith("/model"):
|
||||
assert "/healthy/" in path, "Storage failure must occur before spending on analysis"
|
||||
return httpx.Response(200, json=ModelResult(content='{"observations":[]}', cost=0).model_dump())
|
||||
if path.endswith("/result"):
|
||||
saved.put(Result.model_validate_json(request.content))
|
||||
return httpx.Response(200, json=True)
|
||||
assert path.endswith(("/progress", "/heartbeat")), path
|
||||
return httpx.Response(200, json=True)
|
||||
|
||||
async with httpx.AsyncClient(base_url="https://proxy.test", transport=httpx.MockTransport(handle)) as client:
|
||||
worker: Final = EngineWorker(client)
|
||||
assert await worker.run_once()
|
||||
failed: Final = saved.get_nowait()
|
||||
assert failed.error.startswith("Worker temporary storage failed.")
|
||||
assert not failed.findings
|
||||
assert not tuple(Path("/tmp").glob("lens-trace-*")), "Failed scan left temporary files behind"
|
||||
assert await worker.run_once()
|
||||
recovered: Final = saved.get_nowait()
|
||||
assert recovered.error == "" and recovered.coverage.screened == 1
|
||||
assert not tuple(Path("/tmp").glob("lens-trace-*"))
|
||||
logging.info("Storage-full scan failed clearly; temporary files cleaned; next scan completed")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
|
|
@ -1440,7 +1440,11 @@ def test_jwt_path_enforces_the_user_model_budget_before_returning():
|
|||
|
||||
from litellm.proxy.auth import user_api_key_auth as auth_module
|
||||
|
||||
tree = ast.parse(textwrap.dedent(inspect.getsource(auth_module._user_api_key_auth_builder)))
|
||||
tree = ast.parse(
|
||||
textwrap.dedent(inspect.getsource(auth_module._user_api_key_auth_builder))
|
||||
+ "\n"
|
||||
+ textwrap.dedent(inspect.getsource(auth_module.validate_resolved_virtual_key))
|
||||
)
|
||||
|
||||
def calls_before_each_return(node):
|
||||
seen_check = []
|
||||
|
|
@ -1481,7 +1485,11 @@ def test_every_jwt_branch_carries_the_user_model_budget():
|
|||
|
||||
from litellm.proxy.auth import user_api_key_auth as auth_module
|
||||
|
||||
tree = ast.parse(textwrap.dedent(inspect.getsource(auth_module._user_api_key_auth_builder)))
|
||||
tree = ast.parse(
|
||||
textwrap.dedent(inspect.getsource(auth_module._user_api_key_auth_builder))
|
||||
+ "\n"
|
||||
+ textwrap.dedent(inspect.getsource(auth_module.validate_resolved_virtual_key))
|
||||
)
|
||||
|
||||
assignments = [
|
||||
node
|
||||
|
|
@ -1614,7 +1622,11 @@ def test_zero_cost_models_skip_the_user_budget_check_on_every_path():
|
|||
|
||||
from litellm.proxy.auth import user_api_key_auth as auth_module
|
||||
|
||||
tree = ast.parse(textwrap.dedent(inspect.getsource(auth_module._user_api_key_auth_builder)))
|
||||
tree = ast.parse(
|
||||
textwrap.dedent(inspect.getsource(auth_module._user_api_key_auth_builder))
|
||||
+ "\n"
|
||||
+ textwrap.dedent(inspect.getsource(auth_module.validate_resolved_virtual_key))
|
||||
)
|
||||
|
||||
def guarded_by_skip(node: ast.AST, target: ast.AST) -> bool:
|
||||
for parent in ast.walk(node):
|
||||
|
|
@ -1755,7 +1767,11 @@ def test_mapped_key_jwt_falls_through_to_the_shared_user_budget_attach():
|
|||
|
||||
from litellm.proxy.auth import user_api_key_auth as auth_module
|
||||
|
||||
tree = ast.parse(textwrap.dedent(inspect.getsource(auth_module._user_api_key_auth_builder)))
|
||||
tree = ast.parse(
|
||||
textwrap.dedent(inspect.getsource(auth_module._user_api_key_auth_builder))
|
||||
+ "\n"
|
||||
+ textwrap.dedent(inspect.getsource(auth_module.validate_resolved_virtual_key))
|
||||
)
|
||||
|
||||
# Half one: the shared block copies the user row's budget onto the token.
|
||||
copies_user_row = [
|
||||
|
|
|
|||
|
|
@ -200,3 +200,45 @@ def test_batch_snapshot_keeps_feedback_identity_and_only_current_evidence() -> N
|
|||
assert snapshot.title == "Updated wording"
|
||||
assert snapshot.evidence == draft.evidence
|
||||
assert snapshot.revision == 2
|
||||
|
||||
|
||||
@pytest.mark.parametrize("explicit_reference", (False, True))
|
||||
def test_issue_and_pattern_with_same_title_keep_independent_feedback(explicit_reference: bool) -> None:
|
||||
from litellm.proxy.engine.state import snapshot_finding
|
||||
|
||||
original: Final = engine()
|
||||
issue: Final = merge_finding(original, finding("old"), 1, NOW).model_copy(
|
||||
update={"status": "dismissed", "reason": "Expected retry"}
|
||||
)
|
||||
reviewed: Final = original.model_copy(update={"findings": (issue,)})
|
||||
draft: Final = finding("new").model_copy(
|
||||
update={"kind": "pattern", "existing_finding_id": issue.id if explicit_reference else None}
|
||||
)
|
||||
pattern: Final = merge_finding(reviewed, draft, 1, NOW)
|
||||
assert pattern.id != issue.id
|
||||
assert pattern.kind == "pattern"
|
||||
assert pattern.status == "open" and pattern.reason == ""
|
||||
assert pattern.occurrences == ("new",)
|
||||
assert snapshot_finding(reviewed, draft, 1, NOW).id == pattern.id
|
||||
both: Final = reviewed.model_copy(update={"findings": (issue, pattern)})
|
||||
assert merge_finding(both, finding("again"), 1, NOW).id == issue.id
|
||||
assert merge_finding(both, finding("again"), 1, NOW).status == "dismissed"
|
||||
|
||||
|
||||
def test_legacy_finding_identity_preserves_feedback_only_for_same_kind_and_check() -> None:
|
||||
import hashlib
|
||||
|
||||
original: Final = engine()
|
||||
draft: Final = finding("old")
|
||||
legacy_id: Final = hashlib.sha256(f"{original.id}:{draft.check_id}:{draft.title.lower()}".encode()).hexdigest()[:24]
|
||||
legacy: Final = merge_finding(original, draft, 1, NOW).model_copy(
|
||||
update={"id": legacy_id, "status": "dismissed", "reason": "Accepted"}
|
||||
)
|
||||
reviewed: Final = original.model_copy(update={"findings": (legacy,)})
|
||||
repeated: Final = merge_finding(reviewed, finding("new"), 2, NOW)
|
||||
assert repeated.id == legacy_id
|
||||
assert repeated.status == "dismissed" and repeated.reason == "Accepted"
|
||||
other: Final = finding("new").model_copy(update={"check_id": "different", "existing_finding_id": legacy_id})
|
||||
separate: Final = merge_finding(reviewed, other, 2, NOW)
|
||||
assert separate.id != legacy_id
|
||||
assert separate.status == "open" and separate.reason == ""
|
||||
|
|
|
|||
|
|
@ -0,0 +1,47 @@
|
|||
import { screen } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import { renderWithProviders, testQueryClient } from "@/../tests/test-utils";
|
||||
import { apiClient } from "@/components/networking";
|
||||
import { AnalysisKey } from "./AnalysisKey";
|
||||
|
||||
vi.mock("@/components/networking", () => ({ apiClient: { get: vi.fn(), post: vi.fn() } }));
|
||||
|
||||
describe("Lens billing key", () => {
|
||||
beforeEach(() => {
|
||||
testQueryClient.clear();
|
||||
vi.clearAllMocks();
|
||||
});
|
||||
it("creates a normal key and only passes its ID to worker settings", async () => {
|
||||
const user = userEvent.setup();
|
||||
const changed = vi.fn();
|
||||
vi.mocked(apiClient.get).mockResolvedValue({ keys: [], total_pages: 0 });
|
||||
vi.mocked(apiClient.post).mockResolvedValue({ token_id: "b".repeat(64), key: "sk-secret-not-for-settings" });
|
||||
renderWithProviders(<AnalysisKey accessToken="test" value={null} onChange={changed} name="Research" />);
|
||||
await user.click(screen.getByRole("button", { name: "Create worker key" }));
|
||||
expect(await screen.findByRole("combobox", { name: "Charge analysis to" })).toHaveValue("Lens: Research");
|
||||
expect(apiClient.post).toHaveBeenCalledWith("/key/generate", {
|
||||
accessToken: "test",
|
||||
body: { key_alias: "Lens: Research", models: [], metadata: { purpose: "lens" } },
|
||||
});
|
||||
expect(changed).toHaveBeenCalledExactlyOnceWith("b".repeat(64));
|
||||
expect(screen.queryByText("sk-secret-not-for-settings")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("pages existing keys without dropping the selected billing key", async () => {
|
||||
const user = userEvent.setup();
|
||||
const changed = vi.fn();
|
||||
vi.mocked(apiClient.get).mockImplementation(async (_path, options) => ({
|
||||
keys:
|
||||
options?.query?.page === "2"
|
||||
? [{ token: "c".repeat(64), key_alias: "Second page" }]
|
||||
: [{ token: "a".repeat(64), key_alias: "First page" }],
|
||||
total_pages: 2,
|
||||
}));
|
||||
renderWithProviders(<AnalysisKey accessToken="test" value={null} onChange={changed} name="Research" />);
|
||||
await user.click(screen.getByRole("combobox", { name: "Charge analysis to" }));
|
||||
await user.click(await screen.findByRole("option", { name: "Load more keys" }));
|
||||
await user.click(await screen.findByRole("option", { name: "Second page" }));
|
||||
expect(changed).toHaveBeenCalledExactlyOnceWith("c".repeat(64));
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,144 @@
|
|||
"use client";
|
||||
|
||||
import { useState } from "react";
|
||||
import { useInfiniteQuery } from "@tanstack/react-query";
|
||||
import { z } from "zod";
|
||||
import { apiClient } from "@/components/networking";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import {
|
||||
Combobox,
|
||||
ComboboxContent,
|
||||
ComboboxEmpty,
|
||||
ComboboxInput,
|
||||
ComboboxItem,
|
||||
ComboboxList,
|
||||
} from "@/components/ui/combobox";
|
||||
|
||||
const keySchema = z.object({ token: z.string(), key_alias: z.string().nullable().optional() });
|
||||
const pageSchema = z.object({ keys: z.array(keySchema), total_pages: z.number() });
|
||||
type Key = z.infer<typeof keySchema>;
|
||||
|
||||
export function AnalysisKey({
|
||||
accessToken,
|
||||
value,
|
||||
onChange,
|
||||
name,
|
||||
}: {
|
||||
accessToken: string;
|
||||
value: string | null;
|
||||
onChange: (key: string | null) => void;
|
||||
name: string;
|
||||
}) {
|
||||
const [query, setQuery] = useState("");
|
||||
const [selected, setSelected] = useState<Key | null>(value ? { token: value } : null);
|
||||
const [creating, setCreating] = useState(false);
|
||||
const [error, setError] = useState("");
|
||||
const queryOptions = {
|
||||
queryKey: ["lens-analysis-keys", accessToken, query],
|
||||
initialPageParam: 1,
|
||||
queryFn: async ({ pageParam, signal }: { pageParam: number; signal: AbortSignal }) =>
|
||||
pageSchema.parse(
|
||||
await apiClient.get("/key/list", {
|
||||
accessToken,
|
||||
signal,
|
||||
query: {
|
||||
page: String(pageParam),
|
||||
size: "25",
|
||||
return_full_object: "true",
|
||||
key_alias: query || undefined,
|
||||
substring_matching: "true",
|
||||
include_team_keys: "true",
|
||||
include_created_by_keys: "true",
|
||||
status: "active",
|
||||
},
|
||||
}),
|
||||
),
|
||||
getNextPageParam: (lastPage: z.infer<typeof pageSchema>, pages: z.infer<typeof pageSchema>[]) =>
|
||||
pages.length < lastPage.total_pages ? pages.length + 1 : undefined,
|
||||
};
|
||||
const keyPages = useInfiniteQuery(queryOptions);
|
||||
const keys = keyPages.data?.pages.flatMap((page) => page.keys) ?? [];
|
||||
const choice = keys.find((key) => key.token === value) ?? selected;
|
||||
const loading = keyPages.isFetching;
|
||||
|
||||
const create = async () => {
|
||||
setCreating(true);
|
||||
setError("");
|
||||
try {
|
||||
const result = await apiClient.post("/key/generate", {
|
||||
accessToken,
|
||||
body: {
|
||||
key_alias: `Lens: ${name}`,
|
||||
models: [],
|
||||
metadata: { purpose: "lens" },
|
||||
},
|
||||
});
|
||||
if (!result.token_id) throw new Error("The proxy did not return the new key's ID");
|
||||
const key = { token: result.token_id, key_alias: `Lens: ${name}` };
|
||||
setSelected(key);
|
||||
onChange(key.token);
|
||||
} catch (cause) {
|
||||
setError(cause instanceof Error ? cause.message : "Could not create a key");
|
||||
} finally {
|
||||
setCreating(false);
|
||||
}
|
||||
};
|
||||
const changeKey = (key: Key | null, details: { cancel: () => void }) => {
|
||||
if (key?.token === "load-more") {
|
||||
details.cancel();
|
||||
if (!loading) void keyPages.fetchNextPage();
|
||||
return;
|
||||
}
|
||||
setSelected(key);
|
||||
onChange(key?.token ?? null);
|
||||
};
|
||||
const choices = choice && !keys.some((key) => key.token === choice.token) ? [choice, ...keys] : keys;
|
||||
const items = keyPages.hasNextPage
|
||||
? [...choices, { token: "load-more", key_alias: loading ? "Loading…" : "Load more keys" }]
|
||||
: choices;
|
||||
return (
|
||||
<div className="space-y-2">
|
||||
<p className="text-sm">Charge analysis to</p>
|
||||
<div className="flex items-start gap-2">
|
||||
<div className="min-w-0 flex-1">
|
||||
<Combobox
|
||||
items={items}
|
||||
value={choice}
|
||||
filter={null}
|
||||
itemToStringLabel={(key: Key) => key.key_alias || `${key.token.slice(0, 8)}…`}
|
||||
isItemEqualToValue={(a: Key, b: Key) => a.token === b.token}
|
||||
onInputValueChange={(text, details) => {
|
||||
if (details.reason === "input-change" || details.reason === "input-clear") {
|
||||
setQuery(text);
|
||||
}
|
||||
}}
|
||||
onValueChange={changeKey}
|
||||
>
|
||||
<ComboboxInput aria-label="Charge analysis to" placeholder="Search existing keys" />
|
||||
<ComboboxContent>
|
||||
<ComboboxEmpty>{loading ? "Loading keys…" : "No matching keys"}</ComboboxEmpty>
|
||||
<ComboboxList>
|
||||
{(key: Key) => (
|
||||
<ComboboxItem key={key.token} value={key}>
|
||||
{key.key_alias || `${key.token.slice(0, 8)}…`}
|
||||
</ComboboxItem>
|
||||
)}
|
||||
</ComboboxList>
|
||||
</ComboboxContent>
|
||||
</Combobox>
|
||||
</div>
|
||||
<Button variant="outline" disabled={creating} onClick={() => void create()}>
|
||||
{creating ? "Creating…" : "Create worker key"}
|
||||
</Button>
|
||||
</div>
|
||||
<p className="text-xs text-muted-foreground">
|
||||
Spend appears under this key in API Keys. Its permissions and limits apply.
|
||||
</p>
|
||||
{(error || keyPages.error) && (
|
||||
<p role="alert" className="text-sm text-destructive">
|
||||
{error || keyPages.error?.message}
|
||||
</p>
|
||||
)}
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
|
@ -193,7 +193,14 @@ it("runs saved settings immediately without opening setup", async () => {
|
|||
engines: [engine],
|
||||
tracing_enabled: true,
|
||||
workers: [
|
||||
{ id: "worker", name: "Worker", revoked: false, scope: engine.scope, last_seen: new Date().toISOString() },
|
||||
{
|
||||
id: "worker",
|
||||
name: "Worker",
|
||||
revoked: false,
|
||||
analysis_key_id: "a".repeat(64),
|
||||
scope: engine.scope,
|
||||
last_seen: new Date().toISOString(),
|
||||
},
|
||||
],
|
||||
};
|
||||
if (path === "/engine/lens/runs") return engine.jobs;
|
||||
|
|
|
|||
|
|
@ -95,7 +95,9 @@ export function EngineView({ accessToken, readOnly = false }: { accessToken: str
|
|||
const showEmpty = !query.isLoading && !query.error && engines.length === 0;
|
||||
const engine = engines.find((e) => e.id === selected) ?? engines[0];
|
||||
const connected =
|
||||
query.data?.workers?.some((w) => !w.revoked && query.dataUpdatedAt - Date.parse(w.last_seen) < 120000) ?? false;
|
||||
query.data?.workers?.some(
|
||||
(w) => !w.revoked && w.analysis_key_id && query.dataUpdatedAt - Date.parse(w.last_seen) < 120000,
|
||||
) ?? false;
|
||||
const historyQuery = {
|
||||
queryKey: ["lens-history", engine?.id, historyOffset, accessToken],
|
||||
enabled: !!engine,
|
||||
|
|
|
|||
|
|
@ -1,12 +1,12 @@
|
|||
import { screen } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
import { renderWithProviders } from "@/../tests/test-utils";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import { renderWithProviders, testQueryClient } from "@/../tests/test-utils";
|
||||
import { apiClient } from "@/components/networking";
|
||||
import { WorkerSetup } from "./WorkerSetup";
|
||||
|
||||
vi.mock("@/components/networking", () => ({
|
||||
apiClient: { post: vi.fn() },
|
||||
apiClient: { get: vi.fn(), post: vi.fn(), put: vi.fn() },
|
||||
proxyBaseUrl: "https://gateway.example/proxy",
|
||||
}));
|
||||
|
||||
|
|
@ -18,10 +18,19 @@ const created = {
|
|||
last_seen: "1970-01-01T00:00:00Z",
|
||||
scope: { all_teams: true, api_key_hash: "", team_id: "" },
|
||||
revoked: false,
|
||||
analysis_key_id: "b".repeat(64),
|
||||
},
|
||||
};
|
||||
|
||||
describe("Worker setup", () => {
|
||||
beforeEach(() => {
|
||||
testQueryClient.clear();
|
||||
vi.clearAllMocks();
|
||||
vi.mocked(apiClient.get).mockResolvedValue({
|
||||
keys: [{ token: "b".repeat(64), key_alias: "Analysis" }],
|
||||
total_pages: 1,
|
||||
});
|
||||
});
|
||||
it("generates a complete command using one worker credential and the configured proxy address", async () => {
|
||||
vi.mocked(apiClient.post).mockResolvedValue(created);
|
||||
const user = userEvent.setup();
|
||||
|
|
@ -29,7 +38,14 @@ describe("Worker setup", () => {
|
|||
expect(screen.getByRole("textbox", { name: "Your LiteLLM deployment URL" })).toHaveValue(
|
||||
"https://gateway.example/proxy",
|
||||
);
|
||||
expect(screen.getByRole("button", { name: "Generate setup command" })).toBeDisabled();
|
||||
await user.click(screen.getByRole("combobox", { name: "Charge analysis to" }));
|
||||
await user.click(await screen.findByRole("option", { name: "Analysis" }));
|
||||
await user.click(screen.getByRole("button", { name: "Generate setup command" }));
|
||||
expect(apiClient.post).toHaveBeenCalledWith("/engine/workers/register", {
|
||||
accessToken: "admin",
|
||||
body: { name: "Lens analyzer", analysis_key_id: "b".repeat(64) },
|
||||
});
|
||||
expect(screen.getByRole("status")).toHaveTextContent("Waiting for your analyzer to connect");
|
||||
await user.click(screen.getByRole("button", { name: "Copy Docker command" }));
|
||||
const command = await navigator.clipboard.readText();
|
||||
|
|
@ -38,4 +54,28 @@ describe("Worker setup", () => {
|
|||
expect(command).toContain("--add-host host.docker.internal:host-gateway");
|
||||
expect(command).toContain("ghcr.io/berriai/litellm-lens-worker@sha256:");
|
||||
});
|
||||
it("assigns billing to an existing worker without replacing its access token", async () => {
|
||||
const user = userEvent.setup();
|
||||
const changed = vi.fn();
|
||||
vi.mocked(apiClient.put).mockResolvedValue(created.worker);
|
||||
renderWithProviders(
|
||||
<WorkerSetup
|
||||
accessToken="admin"
|
||||
workers={[{ ...created.worker, analysis_key_id: null }]}
|
||||
onClose={vi.fn()}
|
||||
onChanged={changed}
|
||||
/>,
|
||||
);
|
||||
expect(screen.getByText("Billing key required")).toBeInTheDocument();
|
||||
await user.click(screen.getByRole("button", { name: "Billing key" }));
|
||||
await user.click(screen.getByRole("combobox", { name: "Charge analysis to" }));
|
||||
await user.click(await screen.findByRole("option", { name: "Analysis" }));
|
||||
await user.click(screen.getByRole("button", { name: "Save billing key" }));
|
||||
expect(apiClient.put).toHaveBeenCalledWith("/engine/workers/worker/billing-key", {
|
||||
accessToken: "admin",
|
||||
body: { analysis_key_id: "b".repeat(64) },
|
||||
});
|
||||
expect(changed).toHaveBeenCalledOnce();
|
||||
expect(apiClient.post).not.toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -6,10 +6,11 @@ import { Dialog, DialogContent, DialogHeader, DialogTitle, DialogDescription } f
|
|||
import { Input } from "@/components/ui/input";
|
||||
import { serverRootPath } from "@/lib/serverRootPath";
|
||||
import { apiClient, proxyBaseUrl } from "@/components/networking";
|
||||
import { AnalysisKey } from "./AnalysisKey";
|
||||
import type { EngineList, WorkerCreated } from "./engineData";
|
||||
|
||||
export const LENS_WORKER_IMAGE =
|
||||
"ghcr.io/berriai/litellm-lens-worker@sha256:40fdb82113dd4474cb6e833cf28552487d87c8baf61693a1c3fc2863b7968c6a";
|
||||
"ghcr.io/berriai/litellm-lens-worker@sha256:c41e932eaf3e4efbcaf8cc5027c7e93021e5b2823f21cb8785cd107e37b91c9a";
|
||||
|
||||
function initialProxyAddress(): string {
|
||||
const url = new URL(proxyBaseUrl || serverRootPath, window.location.origin);
|
||||
|
|
@ -21,6 +22,7 @@ export function workerSetupCommand(address: string, token: string): string {
|
|||
const quote = (value: string) => "'" + value.replaceAll("'", "'\\''") + "'";
|
||||
return [
|
||||
"docker run -d --restart unless-stopped --read-only --cap-drop ALL",
|
||||
" --tmpfs /tmp:rw,noexec,nosuid,size=1g",
|
||||
" --security-opt no-new-privileges --platform linux/amd64 --add-host host.docker.internal:host-gateway",
|
||||
` -e ${quote("LITELLM_URL=" + address)}`,
|
||||
` -e ${quote("LENS_WORKER_TOKEN=" + token)}`,
|
||||
|
|
@ -28,6 +30,11 @@ export function workerSetupCommand(address: string, token: string): string {
|
|||
].join(" \\\n");
|
||||
}
|
||||
|
||||
function workerStatus(worker: EngineList["workers"][number], now: number): string {
|
||||
if (!worker.analysis_key_id) return "Billing key required";
|
||||
return now - Date.parse(worker.last_seen) < 120000 ? "Connected · ready to analyze" : "Not connected";
|
||||
}
|
||||
|
||||
export function WorkerSetup({
|
||||
accessToken,
|
||||
workers,
|
||||
|
|
@ -44,19 +51,37 @@ export function WorkerSetup({
|
|||
const timer = window.setInterval(() => setNow(Date.now()), 15000);
|
||||
return () => window.clearInterval(timer);
|
||||
}, []);
|
||||
const [analysisKey, setAnalysisKey] = useState<string | null>(null);
|
||||
const [editingWorker, setEditingWorker] = useState<string | null>(null);
|
||||
const [address, setAddress] = useState(initialProxyAddress);
|
||||
const [copied, setCopied] = useState(false);
|
||||
const [created, setCreated] = useState<WorkerCreated | null>(null);
|
||||
const [error, setError] = useState("");
|
||||
const [busy, setBusy] = useState(false);
|
||||
const actionLabel = editingWorker ? "Save billing key" : "Generate setup command";
|
||||
const editBilling = (worker: EngineList["workers"][number]) => {
|
||||
setCreated(null);
|
||||
setEditingWorker(worker.id);
|
||||
setAnalysisKey(worker.analysis_key_id ?? null);
|
||||
};
|
||||
const createWorker = async () => {
|
||||
setBusy(true);
|
||||
setError("");
|
||||
try {
|
||||
if (editingWorker) {
|
||||
await apiClient.put(`/engine/workers/${editingWorker}/billing-key`, {
|
||||
accessToken,
|
||||
body: { analysis_key_id: analysisKey },
|
||||
});
|
||||
setEditingWorker(null);
|
||||
setAnalysisKey(null);
|
||||
onChanged();
|
||||
return;
|
||||
}
|
||||
setCreated(
|
||||
await apiClient.post<WorkerCreated>("/engine/workers/register", {
|
||||
accessToken,
|
||||
body: { name: "Lens analyzer" },
|
||||
body: { name: "Lens analyzer", analysis_key_id: analysisKey },
|
||||
}),
|
||||
);
|
||||
onChanged();
|
||||
|
|
@ -80,13 +105,31 @@ export function WorkerSetup({
|
|||
Lens reads your agents’ logs and finds issues in the background. Run its analyzer once with Docker.
|
||||
</DialogDescription>
|
||||
</DialogHeader>
|
||||
<label className="grid gap-2 text-sm">
|
||||
Your LiteLLM deployment URL
|
||||
<Input value={address} onChange={(event) => setAddress(event.target.value)} />
|
||||
</label>
|
||||
<p className="text-xs text-muted-foreground">
|
||||
The analyzer connects to this deployment to read logs and save findings.
|
||||
</p>
|
||||
{!editingWorker && (
|
||||
<>
|
||||
<label className="grid gap-2 text-sm">
|
||||
Your LiteLLM deployment URL
|
||||
<Input value={address} onChange={(event) => setAddress(event.target.value)} />
|
||||
</label>
|
||||
<p className="text-xs text-muted-foreground">
|
||||
The analyzer connects to this deployment to read logs and save findings.
|
||||
</p>
|
||||
</>
|
||||
)}
|
||||
{editingWorker && (
|
||||
<p className="text-sm font-medium">
|
||||
Billing for {workers.find((worker) => worker.id === editingWorker)?.name}
|
||||
</p>
|
||||
)}
|
||||
{!created && (
|
||||
<AnalysisKey
|
||||
key={editingWorker ?? "new"}
|
||||
accessToken={accessToken}
|
||||
value={analysisKey}
|
||||
onChange={setAnalysisKey}
|
||||
name="worker"
|
||||
/>
|
||||
)}
|
||||
{created ? (
|
||||
<div className="space-y-3">
|
||||
<p className="text-sm font-medium">Run this command on your server</p>
|
||||
|
|
@ -115,8 +158,19 @@ export function WorkerSetup({
|
|||
</p>
|
||||
</div>
|
||||
) : (
|
||||
<Button disabled={busy || !address.trim()} onClick={createWorker}>
|
||||
{busy ? "Generating…" : "Generate setup command"}
|
||||
<Button disabled={busy || !address.trim() || !analysisKey} onClick={createWorker}>
|
||||
{busy ? "Saving…" : actionLabel}
|
||||
</Button>
|
||||
)}
|
||||
{editingWorker && (
|
||||
<Button
|
||||
variant="ghost"
|
||||
onClick={() => {
|
||||
setEditingWorker(null);
|
||||
setAnalysisKey(null);
|
||||
}}
|
||||
>
|
||||
Cancel
|
||||
</Button>
|
||||
)}
|
||||
{workers
|
||||
|
|
@ -125,10 +179,11 @@ export function WorkerSetup({
|
|||
<div key={worker.id} className="flex justify-between items-center border-t pt-3 text-sm">
|
||||
<span>
|
||||
{worker.name}
|
||||
<span className="block text-xs text-muted-foreground">
|
||||
{now - Date.parse(worker.last_seen) < 120000 ? "Connected · ready to analyze" : "Not connected"}
|
||||
</span>
|
||||
<span className="block text-xs text-muted-foreground">{workerStatus(worker, now)}</span>
|
||||
</span>
|
||||
<Button variant="ghost" size="sm" onClick={() => editBilling(worker)}>
|
||||
Billing key
|
||||
</Button>
|
||||
<Button
|
||||
variant="ghost"
|
||||
size="sm"
|
||||
|
|
|
|||
61
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
61
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -4984,6 +4984,23 @@ export interface paths {
|
|||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/engine/workers/{worker_id}/billing-key": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path?: never;
|
||||
cookie?: never;
|
||||
};
|
||||
get?: never;
|
||||
/** Set Worker Billing */
|
||||
put: operations["set_worker_billing_engine_workers__worker_id__billing_key_put"];
|
||||
post?: never;
|
||||
delete?: never;
|
||||
options?: never;
|
||||
head?: never;
|
||||
patch?: never;
|
||||
trace?: never;
|
||||
};
|
||||
"/engine/{engine_id}": {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
@ -47996,6 +48013,8 @@ export interface components {
|
|||
};
|
||||
/** Worker */
|
||||
Worker: {
|
||||
/** Analysis Key Id */
|
||||
analysis_key_id?: string | null;
|
||||
/** Id */
|
||||
id: string;
|
||||
/**
|
||||
|
|
@ -48012,6 +48031,11 @@ export interface components {
|
|||
revoked: boolean;
|
||||
scope: components["schemas"]["Scope"];
|
||||
};
|
||||
/** WorkerBilling */
|
||||
WorkerBilling: {
|
||||
/** Analysis Key Id */
|
||||
analysis_key_id: string;
|
||||
};
|
||||
/** WorkerCreated */
|
||||
WorkerCreated: {
|
||||
/** Token */
|
||||
|
|
@ -48020,6 +48044,8 @@ export interface components {
|
|||
};
|
||||
/** WorkerName */
|
||||
WorkerName: {
|
||||
/** Analysis Key Id */
|
||||
analysis_key_id: string;
|
||||
/**
|
||||
* Name
|
||||
* @default Lens worker
|
||||
|
|
@ -55841,6 +55867,41 @@ export interface operations {
|
|||
};
|
||||
};
|
||||
};
|
||||
set_worker_billing_engine_workers__worker_id__billing_key_put: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
header?: never;
|
||||
path: {
|
||||
worker_id: string;
|
||||
};
|
||||
cookie?: never;
|
||||
};
|
||||
requestBody: {
|
||||
content: {
|
||||
"application/json": components["schemas"]["WorkerBilling"];
|
||||
};
|
||||
};
|
||||
responses: {
|
||||
/** @description Successful Response */
|
||||
200: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["Worker"];
|
||||
};
|
||||
};
|
||||
/** @description Validation Error */
|
||||
422: {
|
||||
headers: {
|
||||
[name: string]: unknown;
|
||||
};
|
||||
content: {
|
||||
"application/json": components["schemas"]["HTTPValidationError"];
|
||||
};
|
||||
};
|
||||
};
|
||||
};
|
||||
read_engine_engine__engine_id__get: {
|
||||
parameters: {
|
||||
query?: never;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue