diff --git a/.github/workflows/lens-worker.yml b/.github/workflows/lens-worker.yml index 0798425fd76..53334abaf88 100644 --- a/.github/workflows/lens-worker.yml +++ b/.github/workflows/lens-worker.yml @@ -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: diff --git a/.github/workflows/test-postgres.yml b/.github/workflows/test-postgres.yml index 1ffb7f67f16..ccdf6ef3558 100644 --- a/.github/workflows/test-postgres.yml +++ b/.github/workflows/test-postgres.yml @@ -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 diff --git a/deploy/lens/Dockerfile b/deploy/lens/Dockerfile index dc6f61a4d94..360211194e2 100644 --- a/deploy/lens/Dockerfile +++ b/deploy/lens/Dockerfile @@ -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"] diff --git a/deploy/lens/README.md b/deploy/lens/README.md index 22754f89287..d4bcddf8613 100644 --- a/deploy/lens/README.md +++ b/deploy/lens/README.md @@ -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 diff --git a/deploy/lens/compose.yaml b/deploy/lens/compose.yaml index 773ff00113a..0af04814c1e 100644 --- a/deploy/lens/compose.yaml +++ b/deploy/lens/compose.yaml @@ -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] diff --git a/deploy/lens/screenshots/worker-billing-after.png b/deploy/lens/screenshots/worker-billing-after.png new file mode 100644 index 00000000000..cb8b6991036 Binary files /dev/null and b/deploy/lens/screenshots/worker-billing-after.png differ diff --git a/deploy/lens/screenshots/worker-billing-before.png b/deploy/lens/screenshots/worker-billing-before.png new file mode 100644 index 00000000000..093305fb9e7 Binary files /dev/null and b/deploy/lens/screenshots/worker-billing-before.png differ diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index fa4b36a03aa..7b735152065 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -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" } diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 5194f62cf78..7c5f91d9cf2 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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 diff --git a/litellm/proxy/engine/billing.py b/litellm/proxy/engine/billing.py new file mode 100644 index 00000000000..ca625ed0de6 --- /dev/null +++ b/litellm/proxy/engine/billing.py @@ -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)) diff --git a/litellm/proxy/engine/endpoints.py b/litellm/proxy/engine/endpoints.py index 385c6b2ca5d..d085fe7b289 100644 --- a/litellm/proxy/engine/endpoints.py +++ b/litellm/proxy/engine/endpoints.py @@ -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) diff --git a/litellm/proxy/engine/inference.py b/litellm/proxy/engine/inference.py index 36dbe6e6ffd..687f3832a1b 100644 --- a/litellm/proxy/engine/inference.py +++ b/litellm/proxy/engine/inference.py @@ -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) diff --git a/litellm/proxy/engine/models.py b/litellm/proxy/engine/models.py index 01e05e6745e..33e70ff3eca 100644 --- a/litellm/proxy/engine/models.py +++ b/litellm/proxy/engine/models.py @@ -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 diff --git a/litellm/proxy/engine/repository.py b/litellm/proxy/engine/repository.py index e6e9a272e8b..7e3c2f27282 100644 --- a/litellm/proxy/engine/repository.py +++ b/litellm/proxy/engine/repository.py @@ -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""", diff --git a/litellm/proxy/engine/state.py b/litellm/proxy/engine/state.py index e4f25dc47d5..3ca6e881234 100644 --- a/litellm/proxy/engine/state.py +++ b/litellm/proxy/engine/state.py @@ -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( diff --git a/litellm/proxy/engine/worker.py b/litellm/proxy/engine/worker.py index d7de75ed73c..e71c57ce143 100644 --- a/litellm/proxy/engine/worker.py +++ b/litellm/proxy/engine/worker.py @@ -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." ) diff --git a/tests/integration/spend/test_lens_billing.py b/tests/integration/spend/test_lens_billing.py new file mode 100644 index 00000000000..bd9afdfd954 --- /dev/null +++ b/tests/integration/spend/test_lens_billing.py @@ -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", {}) diff --git a/tests/proxy_behavior/lens/test_lifecycle.py b/tests/proxy_behavior/lens/test_lifecycle.py index c8f24bf8e1b..8a9d3873a29 100644 --- a/tests/proxy_behavior/lens/test_lifecycle.py +++ b/tests/proxy_behavior/lens/test_lifecycle.py @@ -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) diff --git a/tests/proxy_behavior/lens/worker_storage_smoke.py b/tests/proxy_behavior/lens/worker_storage_smoke.py new file mode 100644 index 00000000000..8c80915f978 --- /dev/null +++ b/tests/proxy_behavior/lens/worker_storage_smoke.py @@ -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()) diff --git a/tests/unit/proxy/auth/test_user_api_key_auth.py b/tests/unit/proxy/auth/test_user_api_key_auth.py index 9cdac341b1f..08c2f02a83c 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth.py @@ -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 = [ diff --git a/tests/unit/proxy/engine/test_state.py b/tests/unit/proxy/engine/test_state.py index d731b89964c..8d56f4595da 100644 --- a/tests/unit/proxy/engine/test_state.py +++ b/tests/unit/proxy/engine/test_state.py @@ -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 == "" diff --git a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/AnalysisKey.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/AnalysisKey.integration.test.tsx new file mode 100644 index 00000000000..7a0130bd929 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/AnalysisKey.integration.test.tsx @@ -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(); + 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(); + 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)); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/AnalysisKey.tsx b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/AnalysisKey.tsx new file mode 100644 index 00000000000..c26c42f5700 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/AnalysisKey.tsx @@ -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; + +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(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, pages: z.infer[]) => + 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 ( +
+

Charge analysis to

+
+
+ 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} + > + + + {loading ? "Loading keys…" : "No matching keys"} + + {(key: Key) => ( + + {key.key_alias || `${key.token.slice(0, 8)}…`} + + )} + + + +
+ +
+

+ Spend appears under this key in API Keys. Its permissions and limits apply. +

+ {(error || keyPages.error) && ( +

+ {error || keyPages.error?.message} +

+ )} +
+ ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/EngineView.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/EngineView.integration.test.tsx index 86499e5adf0..384dd9284db 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/EngineView.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/EngineView.integration.test.tsx @@ -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; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/EngineView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/EngineView.tsx index 178a7512821..1ad36c17299 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/EngineView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/EngineView.tsx @@ -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, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/WorkerSetup.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/WorkerSetup.integration.test.tsx index d116283f467..8fdeab444d8 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/WorkerSetup.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/WorkerSetup.integration.test.tsx @@ -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( + , + ); + 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(); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/WorkerSetup.tsx b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/WorkerSetup.tsx index 3052518ce37..7800f445a17 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/WorkerSetup.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/WorkerSetup.tsx @@ -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(null); + const [editingWorker, setEditingWorker] = useState(null); const [address, setAddress] = useState(initialProxyAddress); const [copied, setCopied] = useState(false); const [created, setCreated] = useState(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("/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. - -

- The analyzer connects to this deployment to read logs and save findings. -

+ {!editingWorker && ( + <> + +

+ The analyzer connects to this deployment to read logs and save findings. +

+ + )} + {editingWorker && ( +

+ Billing for {workers.find((worker) => worker.id === editingWorker)?.name} +

+ )} + {!created && ( + + )} {created ? (

Run this command on your server

@@ -115,8 +158,19 @@ export function WorkerSetup({

) : ( - + )} + {editingWorker && ( + )} {workers @@ -125,10 +179,11 @@ export function WorkerSetup({
{worker.name} - - {now - Date.parse(worker.last_seen) < 120000 ? "Connected · ready to analyze" : "Not connected"} - + {workerStatus(worker, now)} +