diff --git a/.github/workflows/test-litellm-ui-unit.yml b/.github/workflows/test-litellm-ui-unit.yml index 8f2199017d9..69cbc082d98 100644 --- a/.github/workflows/test-litellm-ui-unit.yml +++ b/.github/workflows/test-litellm-ui-unit.yml @@ -42,6 +42,11 @@ jobs: - name: Install dependencies run: npm ci + - name: Run UI type tests (Vitest) + env: + CI: "true" + run: npm run test:types + - name: Run UI unit tests (Vitest) env: CI: "true" diff --git a/litellm/integrations/arize/_utils.py b/litellm/integrations/arize/_utils.py index 8c494794858..e7e1ab538d5 100644 --- a/litellm/integrations/arize/_utils.py +++ b/litellm/integrations/arize/_utils.py @@ -1,4 +1,5 @@ import json +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final from typing_extensions import override @@ -12,7 +13,7 @@ from litellm.litellm_core_utils.redact_messages import ( should_redact_message_logging, ) from litellm.litellm_core_utils.safe_json_dumps import safe_dumps -from litellm.types.utils import StandardLoggingPayload +from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall, StandardLoggingPayload if TYPE_CHECKING: from opentelemetry.trace import Span @@ -22,6 +23,7 @@ from litellm.integrations._types.open_inference import ( ImageAttributes, MessageAttributes, MessageContentAttributes, + OpenInferenceMimeTypeValues, OpenInferenceSpanKindValues, SpanAttributes, ToolCallAttributes, @@ -480,6 +482,7 @@ def set_attributes(span: "Span", kwargs, response_obj, attributes: type[BaseLLMO response_obj_for_attrs, slp, ) + _safe_emit("mcp tool attrs", _maybe_set_mcp_tool_attrs, span, kwargs, slp, response_obj_for_attrs) def _sanitize_optional_params(optional_params: dict | None) -> dict: @@ -538,9 +541,12 @@ def _set_request_attributes( if optional_params.get("user"): safe_set_attribute(span, "llm.user", optional_params.get("user")) - if response_obj and response_obj.get("id"): + if not hasattr(response_obj, "get"): + return + + if response_obj.get("id"): safe_set_attribute(span, "llm.response.id", response_obj.get("id")) - if response_obj and response_obj.get("model"): + if response_obj.get("model"): safe_set_attribute(span, "llm.response.model", response_obj.get("model")) @@ -588,6 +594,8 @@ def _coerce_response_obj_for_attrs(response_obj): - dicts and Pydantic models that already expose `.get` are returned unchanged (preserves all current behavior, including the Responses API flow which relies on Pydantic attribute access). + - Pydantic models without `.get` (e.g. the MCP SDK's `CallToolResult`, + logged for `call_mcp_tool` spans) are dumped to a dict. - `httpx.Response` and other text-only responses (passthrough routes) are JSON-decoded so the standard extraction paths can read fields like `id`, `model`, and `usage`. On failure the original object is returned @@ -595,6 +603,9 @@ def _coerce_response_obj_for_attrs(response_obj): """ if response_obj is None or hasattr(response_obj, "get"): return response_obj + dumped: Final = _to_plain_dict(response_obj) + if isinstance(dumped, dict): + return dumped text: Final = getattr(response_obj, "text", None) if isinstance(text, str) and text: try: @@ -1058,3 +1069,65 @@ def _parse_passthrough_response(raw_response_obj, coerced_response_obj, kwargs): except Exception: return None return None + + +def _maybe_set_mcp_tool_attrs( + span: "Span", + kwargs: Mapping[str, object], + standard_logging_payload: StandardLoggingPayload | None, + coerced_response_obj: object, +) -> None: + """Render `call_mcp_tool` spans as OpenInference TOOL spans. + + MCP tool calls carry neither `messages` nor `choices`, so the generic + extraction paths leave Input/Output blank. The tool name and arguments live + in `metadata.mcp_tool_call_metadata`; the result is an MCP `CallToolResult` + whose `content` is a list of typed parts. + """ + if standard_logging_payload is None: + return + if standard_logging_payload.get("call_type") != CallTypes.call_mcp_tool.value: + return + + metadata: Final = standard_logging_payload.get("metadata") + mcp_meta: Final[StandardLoggingMCPToolCall | None] = metadata.get("mcp_tool_call_metadata") if metadata else None + if mcp_meta is None: + return + + tool_name: Final = mcp_meta.get("name") or mcp_meta.get("namespaced_tool_name") + if tool_name: + safe_set_attribute(span, SpanAttributes.TOOL_NAME, tool_name) + + if should_redact_message_logging(kwargs): # pyright: ignore[reportArgumentType] # reads, never mutates + return + + arguments: Final[object] = mcp_meta.get("arguments") + if arguments is not None: + safe_set_attribute(span, SpanAttributes.INPUT_VALUE, safe_dumps(arguments)) + safe_set_attribute(span, SpanAttributes.INPUT_MIME_TYPE, OpenInferenceMimeTypeValues.JSON.value) + + _set_mcp_tool_output(span, coerced_response_obj) + + +def _has_only_text_parts(content: object) -> bool: + return not isinstance(content, list) or all(_coerce_text([part]) is not None for part in content) + + +def _set_mcp_tool_output(span: "Span", coerced_response_obj: object) -> None: + if not isinstance(coerced_response_obj, Mapping): + return + + content: Final[object] = coerced_response_obj.get("content") + text: Final[str | None] = _coerce_text(content) + if text and _has_only_text_parts(content): + safe_set_attribute(span, SpanAttributes.OUTPUT_VALUE, text) + safe_set_attribute(span, SpanAttributes.OUTPUT_MIME_TYPE, OpenInferenceMimeTypeValues.TEXT.value) + return + + structured: Final[object] = coerced_response_obj.get("structuredContent") + payload: Final[object] = content if content else structured if structured is not None else content + if payload is None: + return + + safe_set_attribute(span, SpanAttributes.OUTPUT_VALUE, safe_dumps(payload)) + safe_set_attribute(span, SpanAttributes.OUTPUT_MIME_TYPE, OpenInferenceMimeTypeValues.JSON.value) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index c8ca9dcf57e..dc4f17c7b31 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -4031,6 +4031,7 @@ class LitellmDataForBackendLLMCall(TypedDict, total=False): # deliberately tiny value isn't treated as a deployment health signal (see # cooldown_handlers._trigger_cooldown_for_failed_deployment). client_side_timeout: bool + keepalive_seconds: float | None class LitellmMetadataFromRequestHeaders(TypedDict, total=False): diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 66bfc8b81d8..10142a894a1 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -882,6 +882,19 @@ class LiteLLMProxyRequestSetup: return float(stream_timeout_header) return None + @staticmethod + def _get_keepalive_seconds_from_request(headers: Mapping[str, str]) -> float | None: + """ + Get `keepalive_seconds` from the request headers, for clients (e.g. the + Vercel AI SDK) that can set custom headers more easily than extra body + fields. Subject to the same deployment-level allow_client_keepalive_override + gate as the request body field: see _resolve_keepalive_seconds. + """ + keepalive_seconds_header: Final = headers.get("x-litellm-keepalive-seconds", None) + if keepalive_seconds_header is not None: + return float(keepalive_seconds_header) + return None + @staticmethod def _get_num_retries_from_request(headers: dict) -> int | None: """ @@ -1114,6 +1127,10 @@ class LiteLLMProxyRequestSetup: if num_retries is not None: data["num_retries"] = num_retries + keepalive_seconds: Final = LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request(headers) + if keepalive_seconds is not None: + data["keepalive_seconds"] = keepalive_seconds + return data @staticmethod diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index bc980934f9f..baa1579d2c0 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -15,7 +15,7 @@ import threading import time import traceback import warnings -from collections.abc import AsyncGenerator, Callable, Mapping +from collections.abc import AsyncGenerator, AsyncIterator, Callable, Mapping from datetime import datetime, timedelta, timezone from types import MappingProxyType, UnionType from typing import ( @@ -23,6 +23,7 @@ from typing import ( Any, Final, Literal, + NamedTuple, Optional, TypedDict, Union, @@ -7643,6 +7644,200 @@ def _pop_complete_sse_frame(buffer: str) -> tuple[str | None, str]: return buffer[:frame_end], buffer[frame_end:] +_STREAM_KEEPALIVE: Final = object() + +_KEEPALIVE_MIN_SECONDS: Final = 1.0 +_KEEPALIVE_MAX_SECONDS: Final = 300.0 +_EMPTY_MAPPING: Final[Mapping[str, Any]] = MappingProxyType({}) + + +async def _iter_with_keepalive( + aiter: AsyncIterator[Any], + resolve_keepalive_seconds: Callable[[object], float], + keepalive_seconds: float, +) -> AsyncGenerator[Any, None]: + """Wrap `aiter` with idle-gap heartbeats, re-resolving the interval after each + real chunk via `resolve_keepalive_seconds`. A mid-stream router fallback can + swap in a deployment with a different keepalive policy, including one that + newly enables or newly disables heartbeats, partway through the same stream; + re-resolving against each chunk's own identity (rather than trusting the + interval picked before iteration started, or picked the last time it went + inactive) keeps the heartbeat behavior in sync with whichever deployment + actually produced it, in both directions. While the interval is <= 0, no + task is created and no timeout is awaited: a chunk is forwarded the moment + it arrives, at the same cost as a bare `async for`.""" + pending: asyncio.Task[Any] | None = None # rebind-ok: rebound each loop iteration + current_keepalive_seconds = keepalive_seconds # rebind-ok: re-resolved after each chunk + try: + while True: + if current_keepalive_seconds <= 0: + try: + item = await aiter.__anext__() + except StopAsyncIteration: + break + yield item + current_keepalive_seconds = resolve_keepalive_seconds(item) + continue + + if pending is None: + pending = asyncio.create_task(aiter.__anext__()) + done, _ = await asyncio.wait((pending,), timeout=current_keepalive_seconds) + if not done: + yield _STREAM_KEEPALIVE + continue + try: + item = pending.result() + except StopAsyncIteration: + break + finally: + pending = None + yield item + current_keepalive_seconds = resolve_keepalive_seconds(item) + finally: + if pending is not None and not pending.done(): + pending.cancel() + try: + await pending + except asyncio.CancelledError: + pass + + +class _DeploymentKeepaliveConfig(NamedTuple): + keepalive_seconds: Any + allow_client_override: bool + + +def _keepalive_from_deployment_config( + request_data: Mapping[str, Any], response: object +) -> _DeploymentKeepaliveConfig | None: + if llm_router is None: + return None + + hidden: Final = get_hidden_params_dict(response) + model_id: Final = hidden.get("model_id") + if isinstance(model_id, str) and model_id: + deployment: Final = llm_router.get_deployment(model_id=model_id) + # A populated model_id names the specific deployment that served this + # stream. If it no longer resolves (e.g. removed by a config reload + # mid-stream), that's a stale identity, not an absent one: don't fall + # through to guessing via model_name below, since a currently-live + # sibling deployment's config was never what actually served this + # stream. + if deployment is None: + return None + return _DeploymentKeepaliveConfig( + keepalive_seconds=getattr(deployment.litellm_params, "keepalive_seconds", None), + allow_client_override=bool(getattr(deployment.litellm_params, "allow_client_keepalive_override", False)), + ) + + # No model_id at all to pin down which deployment actually served this + # stream: only trust the fallback when every deployment under this + # model_name agrees on both keepalive_seconds and + # allow_client_keepalive_override (including deployments that leave either + # field unset), so a stream never inherits a sibling deployment's policy. + configs: Final = frozenset( + ( + (deployment_dict.get("litellm_params") or _EMPTY_MAPPING).get("keepalive_seconds"), + bool( + (deployment_dict.get("litellm_params") or _EMPTY_MAPPING).get("allow_client_keepalive_override", False) + ), + ) + for deployment_dict in llm_router.get_model_list(model_name=request_data.get("model")) or () + ) + if len(configs) == 1: + keepalive_seconds, allow_client_override = next(iter(configs)) + return _DeploymentKeepaliveConfig( + keepalive_seconds=keepalive_seconds, allow_client_override=allow_client_override + ) + return None + + +def _is_explicit_keepalive_disable(raw: object) -> bool: + if not isinstance(raw, (int, float, str)): + return False + try: + return float(raw) <= 0 + except ValueError: + return False + + +def _resolve_keepalive_seconds(request_data: Mapping[str, Any], response: object = None) -> float: + deployment_config: Final = _keepalive_from_deployment_config(request_data, response) + deployment_raw: Final = deployment_config.keepalive_seconds if deployment_config is not None else None + allow_client_override: Final = deployment_config.allow_client_override if deployment_config is not None else False + + # An operator setting keepalive_seconds: 0 on a deployment is an explicit hard + # disable: an authenticated client must not be able to re-enable heartbeats + # (and the idle-timeout evasion that comes with them) for a deployment the + # operator opted out of, regardless of what the request body asks for. + if _is_explicit_keepalive_disable(deployment_raw): + return 0.0 + + # keepalive_seconds is operator-only unless the deployment explicitly opts in: + # a client can't unilaterally enable heartbeats (and the LB-idle-timeout + # evasion that comes with them) for a deployment that never configured this. + client_supplied: Final = request_data.get("keepalive_seconds") if allow_client_override else None + raw: Final = client_supplied if client_supplied is not None else deployment_raw + try: + value: Final = float(raw) if isinstance(raw, (int, float, str)) else 0.0 + except ValueError: + return 0.0 + if value <= 0: + return 0.0 + clamped: Final = max(_KEEPALIVE_MIN_SECONDS, min(value, _KEEPALIVE_MAX_SECONDS)) + if clamped != value: + verbose_proxy_logger.info( + "keepalive_seconds=%s clamped to %s [min=%s, max=%s]", + value, + clamped, + _KEEPALIVE_MIN_SECONDS, + _KEEPALIVE_MAX_SECONDS, + ) + return clamped + + +_KEEPALIVE_CACHE_TTL_SECONDS: Final = 5.0 + + +def _make_keepalive_resolver(request_data: Mapping[str, Any]) -> Callable[[object], float]: + """Wrap `_resolve_keepalive_seconds` with a memo keyed on the serving + deployment's model_id. The steady-state case (no mid-stream fallback, the + overwhelming majority of streams) sees the same model_id on every chunk, so + this turns the per-chunk cost from a full `llm_router.get_deployment()` + Pydantic rebuild into a cheap hidden-params read once per + `_KEEPALIVE_CACHE_TTL_SECONDS` for that model_id. The cache expires on its + own rather than living for the life of the stream, so an operator's live + config change (disabling keepalive, revoking client override, or removing + the deployment) is observed within a bounded window instead of being able + to be evaded by an already-in-flight stream indefinitely. A missing/empty + model_id can't be trusted as a cache key (see + `_keepalive_from_deployment_config`'s model_name fallback, which reflects + current router state rather than one deployment's fixed identity), so + those chunks always resolve fresh, matching prior behavior exactly. + """ + last_model_id: str | None = None # rebind-ok: memoized identity of the last-resolved chunk + last_value: float = 0.0 # rebind-ok: cached resolution for last_model_id + last_resolved_at: float = float("-inf") # rebind-ok: monotonic timestamp of the last real resolution + + def _resolve(item: object) -> float: + nonlocal last_model_id, last_value, last_resolved_at + model_id = get_hidden_params_dict(item).get("model_id") + now: Final = time.monotonic() + if ( + isinstance(model_id, str) + and model_id + and model_id == last_model_id + and now - last_resolved_at < _KEEPALIVE_CACHE_TTL_SECONDS + ): + return last_value + value: Final = _resolve_keepalive_seconds(request_data, item) + if isinstance(model_id, str) and model_id: + last_model_id, last_value, last_resolved_at = model_id, value, now + return value + + return _resolve + + async def async_data_generator( response, user_api_key_dict: UserAPIKeyAuth, @@ -7691,7 +7886,28 @@ async def async_data_generator( else: stream_iterator = response - async for chunk in stream_iterator: + # A stream can start on a deployment with keepalive off and fall back + # mid-stream to one that enables it: only skip wrapping altogether when + # there's no router to ever fall back through in the first place (in + # which case _resolve_keepalive_seconds can never return non-zero for + # any chunk of this stream), not merely because the first chunk's + # deployment happens to start with it off. + resolve_keepalive_seconds: Final = _make_keepalive_resolver(request_data) + stream_source: Final = ( + _iter_with_keepalive( + stream_iterator.__aiter__(), + resolve_keepalive_seconds, + resolve_keepalive_seconds(response), + ) + if llm_router is not None + else stream_iterator + ) + + async for item in stream_source: + if item is _STREAM_KEEPALIVE: + yield ": ping\n\n" + continue + chunk = cast(Any, item) # cast-ok: sentinel already handled above, item is a real chunk here if needs_per_chunk_hook: ### CALL HOOKS ### - modify outgoing data chunk, _str_so_far = await _apply_streaming_chunk_hooks( diff --git a/litellm/types/router.py b/litellm/types/router.py index b03796fb14f..4f8c133c20b 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -281,6 +281,12 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams): # Deployment budgets max_budget: float | None = None budget_duration: str | None = None + keepalive_seconds: float | None = None + # keepalive_seconds is operator-only by default: a client's request-level + # value is ignored unless the deployment opts in here. Prevents a client + # from unilaterally enabling heartbeats (and the LB-idle-timeout evasion + # that comes with them) for a deployment that never configured them. + allow_client_keepalive_override: bool | None = False use_in_pass_through: bool | None = False use_litellm_proxy: bool | None = False use_chat_completions_api: bool | None = None @@ -457,6 +463,8 @@ class LiteLLMParamsTypedDict(TypedDict, total=False): # deployment budgets max_budget: float | None budget_duration: str | None + keepalive_seconds: float | None + allow_client_keepalive_override: bool | None # per-deployment cooldown override cooldown_time: float | None diff --git a/litellm/types/utils.py b/litellm/types/utils.py index abf9382845b..8c7664257df 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3399,6 +3399,8 @@ all_litellm_params = ( + [ "metadata", "litellm_metadata", + "keepalive_seconds", + "allow_client_keepalive_override", "litellm_trace_id", "litellm_request_debug", "guardrails", diff --git a/terraform/provider/RELEASING.md b/terraform/provider/RELEASING.md index 1dc296f29b8..7b359047e2f 100644 --- a/terraform/provider/RELEASING.md +++ b/terraform/provider/RELEASING.md @@ -106,19 +106,23 @@ Before creating a release: 4. **Land the changes in BerriAI/litellm** - Open a PR to `BerriAI/litellm` updating `terraform/provider/CHANGELOG.md` (and any source changes) and merge it. Note the merge commit SHA; the release workflow takes it as `git_ref` + Open a PR to `BerriAI/litellm` updating `terraform/provider/CHANGELOG.md` (and any source changes) and merge it ### 2. Mirror and Tag via project-releaser The provider source lives at `terraform/provider/` in `BerriAI/litellm`; `BerriAI/terraform-provider-litellm` is a thin release mirror. Do not commit or tag the mirror directly +Normally there is nothing to do here. `BerriAI/project-releaser`'s release pipeline runs the same check on every release except `adhoc`, nightly included: it reads the topmost released heading in `terraform/provider/CHANGELOG.md`, probes the mirror for `v`, and dispatches `Publish Terraform provider` only when the changelog has moved ahead of what the mirror carries. Cutting the version heading in step 1 is therefore what releases the provider, and the next release picks it up, so the wait is a day rather than a week + +Dispatch by hand only for an out-of-band release, or to recover a run that failed: + 1. Go to `BerriAI/project-releaser` > **Actions** > `Publish Terraform provider` 2. Click **Run workflow**: - `git_ref`: full 40-char commit SHA from `BerriAI/litellm` to release from - `provider_version`: the new version without the `v` prefix (e.g. `0.3.0`) - `dry_run`: optional; validates without pushing -3. The workflow rsyncs `terraform/provider/` into the mirror repo, commits, and pushes tag `v` -4. The tag push triggers the mirror's `Release` workflow (goreleaser), which is gated by the `production-release` environment approval + +Automatic or manual, the run waits on the `production-release` approval in `project-releaser`, then rsyncs `terraform/provider/` into the mirror repo, commits, and pushes tag `v`. That approval is the only one in the flow. The tag push triggers the mirror's `Release` workflow (goreleaser), which runs unattended **Important**: - Tags must follow the format: `v..` (e.g., `v0.1.2`, `v1.0.0`) diff --git a/tests/test_litellm/integrations/arize/test_arize_utils.py b/tests/test_litellm/integrations/arize/test_arize_utils.py index 83c3351319a..b02fe35cad0 100644 --- a/tests/test_litellm/integrations/arize/test_arize_utils.py +++ b/tests/test_litellm/integrations/arize/test_arize_utils.py @@ -1193,3 +1193,326 @@ def test_arize_coerce_response_obj_returns_original_on_bad_json(): obj = BadJson() assert _coerce_response_obj_for_attrs(obj) is obj + + +def test_arize_mcp_call_tool_result_does_not_break_attribute_setting(): + """`call_mcp_tool` logs the MCP SDK's `CallToolResult`, a Pydantic model + with no `.get`. It used to raise inside `_set_request_attributes`, aborting + the whole attribute block (input messages, invocation params, outputs).""" + from unittest.mock import MagicMock + + from mcp.types import CallToolResult, TextContent + + span = MagicMock() + kwargs = { + "model": "MCP: get_weather", + "standard_logging_object": { + "model_parameters": {"user": "u-1"}, + "metadata": {}, + "call_type": "call_mcp_tool", + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "mcp"}, + } + response_obj = CallToolResult( + content=[TextContent(type="text", text="sunny, 21C")], isError=False + ) + + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + + span.record_exception.assert_not_called() + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.OPENINFERENCE_SPAN_KIND] == "TOOL" + assert written["llm.request.type"] == "call_mcp_tool" + # Emitted after the old crash point, so absent before the fix. + assert written[SpanAttributes.LLM_INVOCATION_PARAMETERS] == '{"user": "u-1"}' + assert written[SpanAttributes.USER_ID] == "u-1" + + +def test_arize_coerce_response_obj_dumps_pydantic_without_get(): + from mcp.types import CallToolResult, TextContent + + from litellm.integrations.arize._utils import _coerce_response_obj_for_attrs + + result = CallToolResult(content=[TextContent(type="text", text="hi")], isError=False) + coerced = _coerce_response_obj_for_attrs(result) + + assert isinstance(coerced, dict) + assert coerced["isError"] is False + assert coerced["content"][0]["text"] == "hi" + + +def test_arize_request_attributes_survive_uncoercible_response_obj(): + """A response object that is neither dict-like nor coercible (binary + passthrough body, SDK object) must not abort attribute setting.""" + from unittest.mock import MagicMock + + span = MagicMock() + kwargs = { + "model": "gpt-4o", + "standard_logging_object": { + "model_parameters": {}, + "metadata": {}, + "call_type": "completion", + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "openai"}, + } + + class Opaque: + pass + + ArizeLogger.set_arize_attributes(span, kwargs, Opaque()) + + span.record_exception.assert_not_called() + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written["llm.provider"] == "openai" + + +def _mcp_kwargs(mcp_tool_call_metadata=None, **overrides): + return { + "model": "MCP: get_weather", + "standard_logging_object": { + "model_parameters": {}, + "metadata": { + "mcp_tool_call_metadata": mcp_tool_call_metadata + or { + "name": "get_weather", + "arguments": {"city": "Seoul"}, + "namespaced_tool_name": "weather-mcp/get_weather", + } + }, + "call_type": "call_mcp_tool", + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "mcp"}, + **overrides, + } + + +def test_arize_mcp_tool_span_renders_name_input_and_output(): + """`call_mcp_tool` spans have no messages/choices, so Input and Output came + out blank. Render them from mcp_tool_call_metadata + CallToolResult.""" + from unittest.mock import MagicMock + + from mcp.types import CallToolResult, TextContent + + span = MagicMock() + response_obj = CallToolResult( + content=[TextContent(type="text", text="sunny, 21C")], isError=False + ) + + ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), response_obj) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.TOOL_NAME] == "get_weather" + assert written[SpanAttributes.INPUT_VALUE] == '{"city": "Seoul"}' + assert written[SpanAttributes.INPUT_MIME_TYPE] == "application/json" + assert written[SpanAttributes.OUTPUT_VALUE] == "sunny, 21C" + assert written[SpanAttributes.OUTPUT_MIME_TYPE] == "text/plain" + + +def test_arize_mcp_tool_span_serializes_non_text_content(): + """Image/resource results have no text part, so fall back to JSON.""" + from unittest.mock import MagicMock + + from mcp.types import CallToolResult, ImageContent + + span = MagicMock() + response_obj = CallToolResult( + content=[ImageContent(type="image", data="Zm9v", mimeType="image/png")], + isError=False, + ) + + ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), response_obj) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.OUTPUT_MIME_TYPE] == "application/json" + assert "image/png" in written[SpanAttributes.OUTPUT_VALUE] + + +def test_arize_mcp_tool_span_respects_message_redaction(): + """Tool arguments and results are user content. With redaction on, only the + tool name may reach the span.""" + from unittest.mock import MagicMock + + from mcp.types import CallToolResult, TextContent + + span = MagicMock() + response_obj = CallToolResult( + content=[TextContent(type="text", text="SSN 123-45-6789")], isError=False + ) + + ArizeLogger.set_arize_attributes( + span, + _mcp_kwargs(standard_callback_dynamic_params={"turn_off_message_logging": True}), + response_obj, + ) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.TOOL_NAME] == "get_weather" + assert SpanAttributes.INPUT_VALUE not in written + assert SpanAttributes.OUTPUT_VALUE not in written + + +def test_arize_non_mcp_span_gets_no_tool_name(): + """The MCP emitter must not fire on ordinary completions.""" + from unittest.mock import MagicMock + + from litellm.types.utils import Choices, ModelResponse + + span = MagicMock() + kwargs = { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hi"}], + "standard_logging_object": { + "model_parameters": {}, + "metadata": {"mcp_tool_call_metadata": {"name": "get_weather"}}, + "call_type": "completion", + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "openai"}, + } + response_obj = ModelResponse( + choices=[Choices(message={"role": "assistant", "content": "hello"})], + model="gpt-4o", + id="r-1", + ) + + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert SpanAttributes.TOOL_NAME not in written + assert written[SpanAttributes.OUTPUT_VALUE] == "hello" + + +def test_arize_mcp_tool_span_renders_empty_arguments(): + from unittest.mock import MagicMock + + from mcp.types import CallToolResult, TextContent + + span = MagicMock() + kwargs = _mcp_kwargs(mcp_tool_call_metadata={"name": "ping", "arguments": {}}) + response_obj = CallToolResult(content=[TextContent(type="text", text="pong")], isError=False) + + ArizeLogger.set_arize_attributes(span, kwargs, response_obj) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.INPUT_VALUE] == "{}" + assert written[SpanAttributes.INPUT_MIME_TYPE] == "application/json" + + +def test_arize_mcp_tool_span_renders_empty_content(): + from unittest.mock import MagicMock + + from mcp.types import CallToolResult + + span = MagicMock() + response_obj = CallToolResult(content=[], isError=False) + + ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), response_obj) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.OUTPUT_VALUE] == "[]" + assert written[SpanAttributes.OUTPUT_MIME_TYPE] == "application/json" + + +def test_arize_mcp_tool_span_falls_back_to_structured_content(): + from unittest.mock import MagicMock + + from mcp.types import CallToolResult + + span = MagicMock() + response_obj = CallToolResult(content=[], structuredContent={"temp_c": 21}, isError=False) + + ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), response_obj) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.OUTPUT_VALUE] == '{"temp_c": 21}' + assert written[SpanAttributes.OUTPUT_MIME_TYPE] == "application/json" + + +def test_arize_list_mcp_tools_response_does_not_break_attribute_setting(): + from unittest.mock import MagicMock + + span = MagicMock() + kwargs = { + "model": "MCP: list_tools", + "messages": [{"role": "user", "content": "list"}], + "standard_logging_object": { + "model_parameters": {}, + "metadata": {}, + "call_type": "list_mcp_tools", + }, + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "mcp"}, + } + + ArizeLogger.set_arize_attributes(span, kwargs, [{"name": "get_weather"}]) + + span.record_exception.assert_not_called() + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written["llm.input_messages.0.message.content"] == "list" + + +def test_arize_mcp_tool_span_serializes_mixed_text_and_media(): + from unittest.mock import MagicMock + + from mcp.types import CallToolResult, ImageContent, TextContent + + span = MagicMock() + response_obj = CallToolResult( + content=[ + TextContent(type="text", text="see image"), + ImageContent(type="image", data="Zm9v", mimeType="image/png"), + ], + isError=False, + ) + + ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), response_obj) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.OUTPUT_MIME_TYPE] == "application/json" + assert "see image" in written[SpanAttributes.OUTPUT_VALUE] + assert "image/png" in written[SpanAttributes.OUTPUT_VALUE] + + +def test_arize_mcp_tool_span_without_response_object_keeps_name_and_input(): + from unittest.mock import MagicMock + + span = MagicMock() + + ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), None) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.TOOL_NAME] == "get_weather" + assert written[SpanAttributes.INPUT_VALUE] == '{"city": "Seoul"}' + assert SpanAttributes.OUTPUT_VALUE not in written + + +def test_arize_mcp_tool_span_without_content_emits_no_output(): + from unittest.mock import MagicMock + + span = MagicMock() + + ArizeLogger.set_arize_attributes(span, _mcp_kwargs(), {"isError": False}) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert written[SpanAttributes.TOOL_NAME] == "get_weather" + assert SpanAttributes.OUTPUT_VALUE not in written + + +def test_arize_mcp_emitter_is_inert_without_a_standard_logging_object(): + from unittest.mock import MagicMock + + span = MagicMock() + kwargs = { + "model": "MCP: get_weather", + "optional_params": {}, + "litellm_params": {"custom_llm_provider": "mcp"}, + } + + ArizeLogger.set_arize_attributes(span, kwargs, None) + + written = {c.args[0]: c.args[1] for c in span.set_attribute.call_args_list} + assert SpanAttributes.TOOL_NAME not in written diff --git a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py index f7e2d276a2e..15758c595c0 100644 --- a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py +++ b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py @@ -196,9 +196,7 @@ async def test_async_assistants_data_generator_hook_failure_yields_error_chunk( async def _noop_failure(*args, **kwargs): return None - monkeypatch.setattr( - ps.proxy_logging_obj, "async_post_call_streaming_hook", _boom_hook - ) + monkeypatch.setattr(ps.proxy_logging_obj, "async_post_call_streaming_hook", _boom_hook) monkeypatch.setattr(ps.proxy_logging_obj, "post_call_failure_hook", _noop_failure) stream = _FakeAssistantsStream([_simple_chunk()]) @@ -385,9 +383,7 @@ def test_get_streaming_fallback_metadata_no_additional_headers(): def test_get_streaming_fallback_metadata_zero_fallback_count(): stream = _FakeStream( [], - hidden_params={ - "additional_headers": {"x-litellm-attempted-fallbacks": 0} - }, + hidden_params={"additional_headers": {"x-litellm-attempted-fallbacks": 0}}, ) assert _get_streaming_fallback_metadata(stream) == (False, None, []) @@ -558,9 +554,7 @@ async def test_apply_streaming_chunk_hooks_appends_to_str_so_far(monkeypatch): async def _passthrough(*, user_api_key_dict, response, data, str_so_far=None): return response - monkeypatch.setattr( - ps.proxy_logging_obj, "async_post_call_streaming_hook", _passthrough - ) + monkeypatch.setattr(ps.proxy_logging_obj, "async_post_call_streaming_hook", _passthrough) new_chunk, new_str = await _apply_streaming_chunk_hooks( chunk=chunk, @@ -870,9 +864,7 @@ async def test_async_data_generator_mid_stream_exception_yields_error_payload( out.append(line) # First entry is the successful "partial" chunk (bytes), last is the error. - assert any( - isinstance(item, str) and item.startswith('data: {"error":') for item in out - ) + assert any(isinstance(item, str) and item.startswith('data: {"error":') for item in out) # --------------------------------------------------------------------------- @@ -914,3 +906,694 @@ def test_select_data_generator_missing_required_kwarg_raises_type_error(): streaming starts.""" with pytest.raises(TypeError): select_data_generator(response=_async_iter([]), user_api_key_dict=_user_auth()) # type: ignore[call-arg] + + +# --------------------------------------------------------------------------- +# SSE keepalive helpers +# --------------------------------------------------------------------------- + + +from litellm.proxy.proxy_server import ( # noqa: E402 + _iter_with_keepalive, + _keepalive_from_deployment_config, + _make_keepalive_resolver, + _resolve_keepalive_seconds, +) +from litellm.proxy.proxy_server import _KEEPALIVE_MAX_SECONDS, _KEEPALIVE_MIN_SECONDS # noqa: E402 + + +@pytest.mark.asyncio +async def test_iter_with_keepalive_hot_path_no_task_wrapping(): + """When keepalive_seconds <= 0, the generator is a transparent pass-through.""" + chunks = [_simple_chunk(content="a"), _simple_chunk(content="b")] + out = [] + async for item in _iter_with_keepalive(_async_iter(chunks), lambda _: 0, keepalive_seconds=0): + out.append(item) + + assert out == chunks + assert ps._STREAM_KEEPALIVE not in out + + +@pytest.mark.asyncio +async def test_iter_with_keepalive_emits_sentinel_when_stream_stalls(): + """With a short keepalive interval and a stalled upstream, _STREAM_KEEPALIVE + sentinels appear before the delayed chunk arrives. The resolver returns a + constant interval, since this test pins the timing mechanics, not + re-resolution.""" + import asyncio + + async def _slow_stream(): + yield _simple_chunk(content="first") + await asyncio.sleep(0.3) + yield _simple_chunk(content="second") + + items = [] + async for item in _iter_with_keepalive(_slow_stream(), lambda _: 0.05, keepalive_seconds=0.05): + items.append(item) + + sentinels = [i for i in items if i is ps._STREAM_KEEPALIVE] + real_chunks = [i for i in items if i is not ps._STREAM_KEEPALIVE] + + assert len(sentinels) >= 2, f"expected >= 2 sentinels during 0.3s stall; got {len(sentinels)}" + assert len(real_chunks) == 2 + assert real_chunks[0].choices[0].delta.content == "first" + assert real_chunks[1].choices[0].delta.content == "second" + + +@pytest.mark.asyncio +async def test_iter_with_keepalive_cancel_on_early_close(): + """Closing the generator early cancels the in-flight task without raising.""" + import asyncio + + async def _infinite_stream(): + while True: + await asyncio.sleep(10) + yield _simple_chunk() + + gen = _iter_with_keepalive(_infinite_stream(), lambda _: 0.05, keepalive_seconds=0.05) + # Advance once to get the sentinel; then close before the real chunk. + first = await gen.__anext__() + assert first is ps._STREAM_KEEPALIVE + # aclose must not raise, and must drain the cancelled task cleanly. + await gen.aclose() + + +@pytest.mark.asyncio +async def test_iter_with_keepalive_disables_after_fallback_lowers_interval(): + """Greptile P1: a mid-stream router fallback can hand off to a deployment + with a different (or disabled) keepalive policy partway through the same + stream. The interval must be re-resolved against each chunk's own identity, + not the one picked before iteration started, or heartbeats keep using the + pre-fallback deployment's policy for the rest of the stream.""" + import asyncio + + async def _slow_stream(): + yield _simple_chunk(content="first") + await asyncio.sleep(0.3) + yield _simple_chunk(content="second") + + def _resolver(item): + # First chunk resolves under the enabled interval used to start the + # wrapper; every chunk after that resolves as if a fallback disabled it. + return 0.0 if item.choices[0].delta.content == "first" else 999.0 + + items = [] + async for item in _iter_with_keepalive(_slow_stream(), _resolver, keepalive_seconds=0.05): + items.append(item) + + sentinels = [i for i in items if i is ps._STREAM_KEEPALIVE] + real_chunks = [i for i in items if i is not ps._STREAM_KEEPALIVE] + + assert sentinels == [], f"expected no sentinels once the resolver disables keepalive; got {len(sentinels)}" + assert len(real_chunks) == 2 + assert real_chunks[0].choices[0].delta.content == "first" + assert real_chunks[1].choices[0].delta.content == "second" + + +@pytest.mark.asyncio +async def test_iter_with_keepalive_enables_after_fallback_raises_interval(): + """Symmetric case: a mid-stream fallback to a deployment with a *shorter* + keepalive interval must take effect immediately, not stay pinned to the + longer interval the stream started with. The interval used to wait for a + chunk is resolved from the *previous* chunk (the only one seen so far when + that wait begins), so the stall has to follow the fallback chunk rather + than precede it: waiting for "third" is where the shorter interval bites.""" + import asyncio + + async def _slow_stream(): + yield _simple_chunk(content="first") + yield _simple_chunk(content="second") + await asyncio.sleep(0.3) + yield _simple_chunk(content="third") + + def _resolver(item): + # "first" resolves under an interval too long to fire before "second" + # arrives; "second" (the fallback chunk) resolves as if the fallback + # deployment enabled a much shorter interval for everything after it. + return 999.0 if item.choices[0].delta.content == "first" else 0.05 + + items = [] + async for item in _iter_with_keepalive(_slow_stream(), _resolver, keepalive_seconds=999.0): + items.append(item) + + sentinels = [i for i in items if i is ps._STREAM_KEEPALIVE] + real_chunks = [i for i in items if i is not ps._STREAM_KEEPALIVE] + + assert len(sentinels) >= 2, ( + f"expected >= 2 sentinels once the resolver enables a short interval; got {len(sentinels)}" + ) + assert len(real_chunks) == 3 + + +@pytest.mark.asyncio +async def test_iter_with_keepalive_activates_from_a_fully_disabled_start(): + """Greptile P1: a stream can start on a deployment with keepalive off + (keepalive_seconds passed in as 0, not merely a long interval) and fall back + mid-stream to one that enables it. The 0-second start must not be treated as + a one-time decision to skip heartbeats for the rest of the stream: no task + is created while inactive, but every chunk still re-resolves so the fallback + chunk can switch the stream into task-wrapped mode.""" + import asyncio + + async def _slow_stream(): + yield _simple_chunk(content="first") + yield _simple_chunk(content="second") + await asyncio.sleep(0.3) + yield _simple_chunk(content="third") + + def _resolver(item): + # "first" resolves to stay off; "second" (the fallback chunk) resolves + # as if the fallback deployment newly enabled a short interval. + return 0.0 if item.choices[0].delta.content == "first" else 0.05 + + items = [] + async for item in _iter_with_keepalive(_slow_stream(), _resolver, keepalive_seconds=0): + items.append(item) + + sentinels = [i for i in items if i is ps._STREAM_KEEPALIVE] + real_chunks = [i for i in items if i is not ps._STREAM_KEEPALIVE] + + assert len(sentinels) >= 2, ( + f"expected >= 2 sentinels once the resolver activates from a disabled start; got {len(sentinels)}" + ) + assert len(real_chunks) == 3 + + +def test_resolve_keepalive_seconds_client_value_ignored_without_override_permission(monkeypatch): + """keepalive_seconds is operator-only by default: a deployment that hasn't set + allow_client_keepalive_override must not let a client's request-level value + change its behavior at all, since that would let any authenticated client + unilaterally enable heartbeats (and the LB-idle-timeout evasion that comes + with them) for a deployment that never opted in.""" + from unittest.mock import MagicMock + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = 15.0 + deployment.litellm_params.allow_client_keepalive_override = False + + router = MagicMock() + router.get_deployment.return_value = deployment + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {"model_id": "deploy-locked"} + + result = _resolve_keepalive_seconds({"model": "my-model", "keepalive_seconds": 1}, response=response) + assert result == 15.0 + + +def test_resolve_keepalive_seconds_request_value_wins_when_override_allowed(monkeypatch): + from unittest.mock import MagicMock + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = None + deployment.litellm_params.allow_client_keepalive_override = True + + router = MagicMock() + router.get_deployment.return_value = deployment + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {"model_id": "deploy-opt-in"} + + result = _resolve_keepalive_seconds({"model": "my-model", "keepalive_seconds": 30}, response=response) + assert result == 30.0 + + +def test_resolve_keepalive_seconds_explicit_zero_disables_when_override_allowed(monkeypatch): + from unittest.mock import MagicMock + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = 20.0 + deployment.litellm_params.allow_client_keepalive_override = True + + router = MagicMock() + router.get_deployment.return_value = deployment + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {"model_id": "deploy-opt-in"} + + result = _resolve_keepalive_seconds({"model": "my-model", "keepalive_seconds": 0}, response=response) + assert result == 0.0 + + +def test_resolve_keepalive_seconds_clamps_below_minimum(monkeypatch): + from unittest.mock import MagicMock + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = None + deployment.litellm_params.allow_client_keepalive_override = True + + router = MagicMock() + router.get_deployment.return_value = deployment + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {"model_id": "deploy-opt-in"} + + result = _resolve_keepalive_seconds({"model": "my-model", "keepalive_seconds": 0.001}, response=response) + assert result == _KEEPALIVE_MIN_SECONDS + + +def test_resolve_keepalive_seconds_clamps_above_maximum(monkeypatch): + from unittest.mock import MagicMock + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = None + deployment.litellm_params.allow_client_keepalive_override = True + + router = MagicMock() + router.get_deployment.return_value = deployment + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {"model_id": "deploy-opt-in"} + + result = _resolve_keepalive_seconds({"model": "my-model", "keepalive_seconds": 9999}, response=response) + assert result == _KEEPALIVE_MAX_SECONDS + + +def test_resolve_keepalive_seconds_non_numeric_returns_zero(monkeypatch): + from unittest.mock import MagicMock + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = None + deployment.litellm_params.allow_client_keepalive_override = True + + router = MagicMock() + router.get_deployment.return_value = deployment + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {"model_id": "deploy-opt-in"} + + result = _resolve_keepalive_seconds({"model": "my-model", "keepalive_seconds": "not-a-number"}, response=response) + assert result == 0.0 + + +def test_resolve_keepalive_seconds_absent_returns_zero(monkeypatch): + monkeypatch.setattr(ps, "llm_router", None) + result = _resolve_keepalive_seconds({}, response=None) + assert result == 0.0 + + +def test_resolve_keepalive_seconds_deployment_disable_cannot_be_overridden_by_request(monkeypatch): + """A deployment that explicitly sets keepalive_seconds: 0 is a hard operator + disable: an authenticated client must not be able to re-enable heartbeats for + that deployment by passing a positive value in the request body, since that + would let a client evade the deployment's idle-timeout behavior at will. This + holds even if the deployment also grants override permission, since an + explicit disable is a stronger, unconditional signal than an override grant.""" + from unittest.mock import MagicMock + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = 0 + deployment.litellm_params.allow_client_keepalive_override = True + + router = MagicMock() + router.get_deployment.return_value = deployment + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {"model_id": "deploy-disabled"} + + result = _resolve_keepalive_seconds({"model": "my-model", "keepalive_seconds": 250}, response=response) + assert result == 0.0 + + +def test_keepalive_from_deployment_config_reads_by_model_id(monkeypatch): + from unittest.mock import MagicMock + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = 45.0 + deployment.litellm_params.allow_client_keepalive_override = True + + router = MagicMock() + router.get_deployment.return_value = deployment + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {"model_id": "deploy-abc"} + + result = _keepalive_from_deployment_config({"model": "my-model"}, response) + assert result == ps._DeploymentKeepaliveConfig(keepalive_seconds=45.0, allow_client_override=True) + router.get_deployment.assert_called_once_with(model_id="deploy-abc") + + +def test_keepalive_from_deployment_config_stale_model_id_does_not_fall_through(monkeypatch): + """A populated model_id names the specific deployment that served the stream. + If that ID no longer resolves (e.g. removed by a config reload mid-stream), + that's a stale identity, not an absent one: it must not fall through to the + model_name fallback, since a currently-live sibling deployment's config was + never what actually served this stream, even if that sibling's config is + unambiguous on its own.""" + from unittest.mock import MagicMock + + router = MagicMock() + router.get_deployment.return_value = None + router.get_model_list.return_value = [ + {"litellm_params": {"keepalive_seconds": 20.0}}, + ] + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {"model_id": "stale-deploy-id"} + + result = _keepalive_from_deployment_config({"model": "slow-model"}, response) + assert result is None + router.get_model_list.assert_not_called() + + +def test_keepalive_from_deployment_config_fallback_by_name(monkeypatch): + from unittest.mock import MagicMock + + router = MagicMock() + router.get_deployment.return_value = None + router.get_model_list.return_value = [ + {"litellm_params": {"keepalive_seconds": 20.0}}, + ] + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {} + + result = _keepalive_from_deployment_config({"model": "slow-model"}, response) + assert result == ps._DeploymentKeepaliveConfig(keepalive_seconds=20.0, allow_client_override=False) + router.get_model_list.assert_called_once_with(model_name="slow-model") + + +def test_keepalive_from_deployment_config_fallback_by_name_agreeing_deployments(monkeypatch): + """Multiple deployments under the same model_name with the same keepalive_seconds + is unambiguous, so the shared value is used even without a model_id.""" + from unittest.mock import MagicMock + + router = MagicMock() + router.get_deployment.return_value = None + router.get_model_list.return_value = [ + {"litellm_params": {"keepalive_seconds": 20.0, "allow_client_keepalive_override": True}}, + {"litellm_params": {"keepalive_seconds": 20.0, "allow_client_keepalive_override": True}}, + ] + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {} + + result = _keepalive_from_deployment_config({"model": "slow-model"}, response) + assert result == ps._DeploymentKeepaliveConfig(keepalive_seconds=20.0, allow_client_override=True) + + +def test_keepalive_from_deployment_config_fallback_by_name_conflicting_deployments(monkeypatch): + """Without a model_id, if deployments under the same model_name disagree on + keepalive_seconds, we can't tell which one served the stream: don't guess and + apply the wrong deployment's interval (or override an explicit disable).""" + from unittest.mock import MagicMock + + router = MagicMock() + router.get_deployment.return_value = None + router.get_model_list.return_value = [ + {"litellm_params": {"keepalive_seconds": 20.0}}, + {"litellm_params": {"keepalive_seconds": 0}}, + ] + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {} + + result = _keepalive_from_deployment_config({"model": "slow-model"}, response) + assert result is None + + +def test_keepalive_from_deployment_config_fallback_by_name_configured_plus_unset(monkeypatch): + """A deployment that leaves keepalive_seconds unset entirely (not explicitly 0) + must not inherit a sibling deployment's configured interval: without a model_id + we can't tell which deployment served the stream, so mixing a configured + deployment with an unconfigured one is just as ambiguous as two conflicting + configured values.""" + from unittest.mock import MagicMock + + router = MagicMock() + router.get_deployment.return_value = None + router.get_model_list.return_value = [ + {"litellm_params": {"keepalive_seconds": 20.0}}, + {"litellm_params": {}}, + ] + + monkeypatch.setattr(ps, "llm_router", router) + + response = MagicMock() + response._hidden_params = {} + + result = _keepalive_from_deployment_config({"model": "slow-model"}, response) + assert result is None + + +def test_keepalive_from_deployment_config_no_router_returns_none(monkeypatch): + monkeypatch.setattr(ps, "llm_router", None) + result = _keepalive_from_deployment_config({"model": "gpt-4"}, None) + assert result is None + + +def test_make_keepalive_resolver_caches_by_model_id(monkeypatch): + """The steady-state case (no fallback): every chunk shares the same + model_id, so the deployment lookup must happen once, not once per chunk.""" + from unittest.mock import MagicMock + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = 5.0 + deployment.litellm_params.allow_client_keepalive_override = False + + router = MagicMock() + router.get_deployment.return_value = deployment + monkeypatch.setattr(ps, "llm_router", router) + + resolve = _make_keepalive_resolver({"model": "my-model"}) + + first = _simple_chunk(content="a") + first._hidden_params = {"model_id": "deploy-steady"} + second = _simple_chunk(content="b") + second._hidden_params = {"model_id": "deploy-steady"} + + assert resolve(first) == 5.0 + assert resolve(second) == 5.0 + router.get_deployment.assert_called_once_with(model_id="deploy-steady") + + +def test_make_keepalive_resolver_reresolves_on_model_id_change(monkeypatch): + """A mid-stream fallback changes model_id: the cache must miss and + re-resolve against the new deployment, not keep serving the stale value.""" + from unittest.mock import MagicMock + + before = MagicMock() + before.litellm_params.keepalive_seconds = 5.0 + before.litellm_params.allow_client_keepalive_override = False + + after = MagicMock() + after.litellm_params.keepalive_seconds = 30.0 + after.litellm_params.allow_client_keepalive_override = False + + router = MagicMock() + router.get_deployment.side_effect = lambda model_id: {"deploy-a": before, "deploy-b": after}[model_id] + monkeypatch.setattr(ps, "llm_router", router) + + resolve = _make_keepalive_resolver({"model": "my-model"}) + + chunk_a = _simple_chunk(content="a") + chunk_a._hidden_params = {"model_id": "deploy-a"} + chunk_b = _simple_chunk(content="b") + chunk_b._hidden_params = {"model_id": "deploy-b"} + + assert resolve(chunk_a) == 5.0 + assert resolve(chunk_b) == 30.0 + assert router.get_deployment.call_count == 2 + + +def test_make_keepalive_resolver_missing_model_id_never_cached(monkeypatch): + """Without a model_id there's no reliable cache key (see the model_name + fallback in _keepalive_from_deployment_config), so every chunk must + re-resolve fresh rather than reuse a stale guess.""" + from unittest.mock import MagicMock + + router = MagicMock() + router.get_deployment.return_value = None + router.get_model_list.return_value = [{"litellm_params": {"keepalive_seconds": 12.0}}] + monkeypatch.setattr(ps, "llm_router", router) + + resolve = _make_keepalive_resolver({"model": "slow-model"}) + + chunk_a = _simple_chunk(content="a") + chunk_a._hidden_params = {} + chunk_b = _simple_chunk(content="b") + chunk_b._hidden_params = {} + + assert resolve(chunk_a) == 12.0 + assert resolve(chunk_b) == 12.0 + assert router.get_model_list.call_count == 2 + + +def test_make_keepalive_resolver_expires_cache_after_ttl(monkeypatch): + """An operator's live config change (revoking override, disabling + keepalive, removing the deployment) must be observed within + _KEEPALIVE_CACHE_TTL_SECONDS, not frozen for the rest of an + already-in-flight stream just because the model_id hasn't changed.""" + from unittest.mock import MagicMock + + before = MagicMock() + before.litellm_params.keepalive_seconds = 20.0 + before.litellm_params.allow_client_keepalive_override = False + + after = MagicMock() + after.litellm_params.keepalive_seconds = 0 + after.litellm_params.allow_client_keepalive_override = False + + router = MagicMock() + router.get_deployment.return_value = before + monkeypatch.setattr(ps, "llm_router", router) + + clock = {"t": 0.0} + monkeypatch.setattr(ps.time, "monotonic", lambda: clock["t"]) + + resolve = _make_keepalive_resolver({"model": "my-model"}) + + chunk = _simple_chunk(content="a") + chunk._hidden_params = {"model_id": "deploy-live"} + + assert resolve(chunk) == 20.0 + assert router.get_deployment.call_count == 1 + + # Still within the TTL: same model_id, cached value reused even though + # the router's live config has since changed underneath it. + router.get_deployment.return_value = after + clock["t"] = ps._KEEPALIVE_CACHE_TTL_SECONDS - 0.01 + assert resolve(chunk) == 20.0 + assert router.get_deployment.call_count == 1 + + # Past the TTL: the config-reload disable is now observed. + clock["t"] = ps._KEEPALIVE_CACHE_TTL_SECONDS + 0.01 + assert resolve(chunk) == 0.0 + assert router.get_deployment.call_count == 2 + + +def test_keepalive_seconds_in_all_litellm_params(): + from litellm.types.utils import all_litellm_params + + assert "keepalive_seconds" in all_litellm_params + + +def test_allow_client_keepalive_override_in_all_litellm_params(): + """allow_client_keepalive_override is a deployment-only control flag: if it's + missing from all_litellm_params, it leaks straight through into the actual + provider API call as an unrecognized field and gets rejected (confirmed live + against the real Anthropic API, which returns 'Extra inputs are not + permitted').""" + from litellm.types.utils import all_litellm_params + + assert "allow_client_keepalive_override" in all_litellm_params + + +@pytest.mark.asyncio +async def test_async_data_generator_emits_ping_heartbeat(monkeypatch): + """When keepalive_seconds is set on a deployment that allows client override, + ': ping' frames appear during upstream stalls.""" + import asyncio + from unittest.mock import MagicMock + + _patch_logging_flags(monkeypatch) + monkeypatch.setattr(ps, "_KEEPALIVE_MIN_SECONDS", 0.05) + + router = MagicMock() + router.get_deployment.return_value = None + router.get_model_list.return_value = [{"litellm_params": {"allow_client_keepalive_override": True}}] + monkeypatch.setattr(ps, "llm_router", router) + + async def _slow_response(): + yield _simple_chunk(content="hello") + await asyncio.sleep(0.4) + yield _simple_chunk(content="world") + + out = [] + async for line in async_data_generator( + response=_slow_response(), + user_api_key_dict=_user_auth(), + request_data={"model": "gpt-4", "keepalive_seconds": 0.05}, + ): + out.append(line) + + pings = [item for item in out if item == ": ping\n\n"] + assert len(pings) >= 2, f"expected >= 2 ping frames; got {len(pings)}" + assert out[-1] == "data: [DONE]\n\n" + + +@pytest.mark.asyncio +async def test_async_data_generator_no_keepalive_no_pings(monkeypatch): + """Without keepalive_seconds, no ': ping' frames are emitted.""" + _patch_logging_flags(monkeypatch) + + out = [] + async for line in async_data_generator( + response=_async_iter([_simple_chunk(content="hello")]), + user_api_key_dict=_user_auth(), + request_data={"model": "gpt-4"}, + ): + out.append(line) + + assert ": ping\n\n" not in out + assert out[-1] == "data: [DONE]\n\n" + + +@pytest.mark.asyncio +async def test_async_data_generator_resolves_deployment_once_per_steady_stream(monkeypatch): + """Regression test for the per-chunk resolver cost: a stream where every + real chunk comes from the same deployment (the common, no-fallback case) + must only pay for one `llm_router.get_deployment()` call, not one per + chunk. Before caching, this asserted 1 but got len(chunks) since the + resolver re-ran the full deployment lookup after every single chunk. + + The very first resolve happens on the raw `response` object before any + chunk is yielded; a bare async generator (unlike the real + CustomStreamWrapper this stands in for) can't carry `_hidden_params`, so + that one call goes through the model_name fallback instead of + `get_deployment` — hence it's asserted separately. + """ + from unittest.mock import MagicMock + + _patch_logging_flags(monkeypatch) + + deployment = MagicMock() + deployment.litellm_params.keepalive_seconds = None + deployment.litellm_params.allow_client_keepalive_override = False + + router = MagicMock() + router.get_deployment.return_value = deployment + router.get_model_list.return_value = [{"litellm_params": {}}] + monkeypatch.setattr(ps, "llm_router", router) + + async def _steady_response(): + for content in ("a", "b", "c", "d", "e"): + chunk = _simple_chunk(content=content) + chunk._hidden_params = {"model_id": "deploy-steady"} + yield chunk + + out = [] + async for line in async_data_generator( + response=_steady_response(), + user_api_key_dict=_user_auth(), + request_data={"model": "gpt-4"}, + ): + out.append(line) + + assert router.get_deployment.call_count == 1 + assert router.get_model_list.call_count == 1 + assert out[-1] == "data: [DONE]\n\n" diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 803094e8d54..f48e1dba601 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -2076,6 +2076,55 @@ def test_get_num_retries_from_request(): assert result == -1 +def test_get_keepalive_seconds_from_request(): + """ + Test LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request method + """ + # Header present with valid float string + headers_with_keepalive = {"x-litellm-keepalive-seconds": "15"} + result = LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request( + headers_with_keepalive + ) + assert result == 15.0 + + # Header not present + result = LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request( + {"Content-Type": "application/json"} + ) + assert result is None + + # Empty headers dictionary + result = LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request({}) + assert result is None + + # Header present with a fractional value + result = LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request( + {"x-litellm-keepalive-seconds": "1.5"} + ) + assert result == 1.5 + + # Header present with invalid value raises ValueError, matching the other + # x-litellm-* numeric header helpers (_get_timeout_from_request, etc.) + with pytest.raises(ValueError): + LiteLLMProxyRequestSetup._get_keepalive_seconds_from_request( + {"x-litellm-keepalive-seconds": "not-a-number"} + ) + + +def test_add_litellm_data_for_backend_llm_call_merges_keepalive_seconds_header(): + """ + The x-litellm-keepalive-seconds header must be merged into the data dict + that add_litellm_data_to_request later data.update()s onto the request body, + the same way x-litellm-timeout/x-litellm-num-retries already are. + """ + result = LiteLLMProxyRequestSetup.add_litellm_data_for_backend_llm_call( + headers={"x-litellm-keepalive-seconds": "20"}, + request_data={}, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + ) + assert result.get("keepalive_seconds") == 20.0 + + def test_add_user_api_key_auth_to_request_metadata(): """ Test that add_user_api_key_auth_to_request_metadata properly adds user API key authentication data to request metadata diff --git a/ui/litellm-dashboard/knip.json b/ui/litellm-dashboard/knip.json index 48b39e8122d..f6cd8ace112 100644 --- a/ui/litellm-dashboard/knip.json +++ b/ui/litellm-dashboard/knip.json @@ -1,6 +1,6 @@ { "$schema": "https://unpkg.com/knip@5/schema.json", - "entry": ["scripts/**/*.{ts,mjs}", "src/components/ui/**/*.{ts,tsx}"], + "entry": ["scripts/**/*.{ts,mjs}", "src/components/ui/**/*.{ts,tsx}", "src/**/*.test-d.{ts,tsx}"], "project": ["src/**/*.{ts,tsx}", "tests/**/*.{ts,tsx}", "scripts/**/*.{ts,mjs}"], "ignore": ["src/lib/http/schema.d.ts"], "ignoreDependencies": [ diff --git a/ui/litellm-dashboard/package.json b/ui/litellm-dashboard/package.json index 3953b41a2c2..62bf5fff4b4 100644 --- a/ui/litellm-dashboard/package.json +++ b/ui/litellm-dashboard/package.json @@ -10,6 +10,7 @@ "lint": "eslint .", "test": "vitest", "test:dot": "vitest --reporter=dot", + "test:types": "vitest --run --typecheck.only", "test:watch": "vitest -w", "test:coverage": "vitest run --coverage", "format": "prettier --write .", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.test.ts index 66dfc43cebb..fa3f15124cf 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.test.ts @@ -820,6 +820,18 @@ describe("useAllTeams", () => { }); const requestedPage = (url: string) => new URLSearchParams(url.split("?")[1]).get("page"); + const requestedUserId = (url: string) => new URLSearchParams(url.split("?")[1]).get("user_id"); + const asRole = (userRole: string, userId = "test-user-id") => + mockUseAuthorized.mockReturnValue({ + accessToken: "test-access-token", + userId, + userRole, + token: "test-token", + userEmail: "test@example.com", + premiumUser: false, + disabledPersonalKeyCreation: null, + showSSOBanner: false, + }); it("paginates /v2/team/list to completion and concatenates every page", async () => { fetchMock.mockImplementation((url: string) => @@ -892,4 +904,63 @@ describe("useAllTeams", () => { await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(2)); }); + + it("scopes the request to the caller for an internal user and returns their teams", async () => { + asRole("Internal User", "member-7"); + fetchMock.mockResolvedValue(pageResponse(mockTeams, 1, 1)); + + const { result } = renderHook(() => useAllTeams(), { wrapper }); + + await waitFor(() => expect(result.current.isSuccess).toBe(true)); + + expect(result.current.data).toEqual(mockTeams); + expect(result.current.data?.length).toBeGreaterThan(0); + expect(requestedUserId(fetchMock.mock.calls[0][0] as string)).toBe("member-7"); + }); + + it("carries user_id on every page of a scoped multi-page result", async () => { + asRole("Internal Viewer", "member-7"); + fetchMock.mockImplementation((url: string) => + Promise.resolve( + requestedPage(url) === "1" ? pageResponse([mockTeams[0]], 1, 2) : pageResponse([mockTeams[1]], 2, 2), + ), + ); + + const { result } = renderHook(() => useAllTeams(), { wrapper }); + + await waitFor(() => expect(result.current.isSuccess).toBe(true)); + + expect(fetchMock).toHaveBeenCalledTimes(2); + const scopes = fetchMock.mock.calls.map((call) => requestedUserId(call[0] as string)); + expect(scopes).toEqual(["member-7", "member-7"]); + }); + + it.each(["Admin", "Admin Viewer", "Org Admin"])( + "sends no user_id for %s so the broad list is left intact", + async (userRole) => { + asRole(userRole); + fetchMock.mockResolvedValue(pageResponse(mockTeams, 1, 1)); + + const { result } = renderHook(() => useAllTeams(), { wrapper }); + + await waitFor(() => expect(result.current.isSuccess).toBe(true)); + + expect(requestedUserId(fetchMock.mock.calls[0][0] as string)).toBeNull(); + }, + ); + + it("refetches when the scope changes even though the access token has not", async () => { + asRole("Internal User", "member-7"); + fetchMock.mockResolvedValue(pageResponse(mockTeams, 1, 1)); + + const { result, rerender } = renderHook(() => useAllTeams(), { wrapper }); + await waitFor(() => expect(result.current.isSuccess).toBe(true)); + expect(fetchMock).toHaveBeenCalledTimes(1); + + asRole("Internal User", "member-8"); + rerender(); + + await waitFor(() => expect(fetchMock).toHaveBeenCalledTimes(2)); + expect(requestedUserId(fetchMock.mock.calls[1][0] as string)).toBe("member-8"); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.ts index 4061026b94d..e209a1d7273 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/teams/useTeams.ts @@ -5,6 +5,7 @@ import { fetchTeams } from "@/app/(dashboard)/networking"; import { createQueryKeys } from "@/app/(dashboard)/hooks/common/queryKeysFactory"; import { teamInfoCall } from "@/components/networking"; import { getProxyBaseUrl, getGlobalLitellmHeaderName, deriveErrorMessage, handleError } from "@/components/networking"; +import { teamListScopeUserId } from "@/utils/roles"; export interface TeamsResponse { teams: Team[]; @@ -116,24 +117,30 @@ export const useTeams = (): UseQueryResult => { const ALL_TEAMS_PAGE_SIZE = 100; -const fetchAllTeamsPaged = async (accessToken: string): Promise => { - const firstPage: TeamsResponse = await teamListCall(accessToken, 1, ALL_TEAMS_PAGE_SIZE); +const fetchAllTeamsPaged = async (accessToken: string, userID: string | null): Promise => { + const firstPage: TeamsResponse = await teamListCall(accessToken, 1, ALL_TEAMS_PAGE_SIZE, { userID }); const totalPages = firstPage.total_pages ?? 1; if (totalPages <= 1) return firstPage.teams; const remainingPages: TeamsResponse[] = await Promise.all( - Array.from({ length: totalPages - 1 }, (_, i) => teamListCall(accessToken, i + 2, ALL_TEAMS_PAGE_SIZE)), + Array.from({ length: totalPages - 1 }, (_, i) => teamListCall(accessToken, i + 2, ALL_TEAMS_PAGE_SIZE, { userID })), ); return [firstPage, ...remainingPages].flatMap((page) => page.teams); }; export const useAllTeams = (): UseQueryResult => { - const { accessToken } = useAuthorized(); + const { accessToken, userId, userRole } = useAuthorized(); + const scopedUserID = teamListScopeUserId(userRole, userId); return useQuery({ queryKey: teamKeys.list({ - filters: { scope: "all", pageSize: ALL_TEAMS_PAGE_SIZE, accessToken: accessToken ?? "" }, + filters: { + scope: "all", + pageSize: ALL_TEAMS_PAGE_SIZE, + accessToken: accessToken ?? "", + userID: scopedUserID ?? "", + }, }), - queryFn: async () => await fetchAllTeamsPaged(accessToken!), + queryFn: async () => await fetchAllTeamsPaged(accessToken!, scopedUserID), enabled: Boolean(accessToken), staleTime: 30000, }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.test.tsx index e3db50b7300..0e4455ee912 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.test.tsx @@ -1,6 +1,6 @@ import React from "react"; import { describe, it, expect, vi, beforeEach } from "vitest"; -import { screen, waitFor, within } from "@testing-library/react"; +import { act, screen, waitFor, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { renderWithProviders } from "../../../../../tests/test-utils"; import UsagePage from "./usage"; @@ -49,6 +49,14 @@ const renderUsage = (overrides: Partial> />, ); +// Width of this window is guarded by "proves the flush window is wide enough". +const flushPendingRequests = async () => { + await act(async () => { + await new Promise((resolve) => setTimeout(resolve, 0)); + await new Promise((resolve) => setTimeout(resolve, 0)); + }); +}; + beforeEach(() => { vi.clearAllMocks(); networking.getProxyUISettings.mockResolvedValue(UNLIMITED_SETTINGS); @@ -185,18 +193,65 @@ describe("old usage page", () => { }); }); - describe("as a non-admin", () => { - it("renders only the All Up tab and skips admin-only queries", async () => { - renderUsage({ userRole: "Internal User" }); + // org_admin is an organization membership role; those users reach the UI as "Internal User". + describe.each(["Internal User", "Internal Viewer", "internal_user", "internal_user_viewer", "Org Admin"])( + "as %s", + (userRole) => { + it("shows the admin-only notice instead of the usage dashboard", async () => { + renderUsage({ userRole }); + + expect(await screen.findByText(/Proxy-wide usage is only available to admin users/i)).toBeInTheDocument(); + expect(screen.queryByRole("tab", { name: "All Up" })).not.toBeInTheDocument(); + }); + + it("fires no /global/spend or /global/activity request", async () => { + renderUsage({ userRole }); + + await screen.findByText(/Proxy-wide usage is only available to admin users/i); + await flushPendingRequests(); + + expect(networking.getProxyUISettings).not.toHaveBeenCalled(); + expect(networking.adminSpendLogsCall).not.toHaveBeenCalled(); + expect(networking.adminTopKeysCall).not.toHaveBeenCalled(); + expect(networking.adminTopModelsCall).not.toHaveBeenCalled(); + expect(networking.adminTopEndUsersCall).not.toHaveBeenCalled(); + expect(networking.teamSpendLogsCall).not.toHaveBeenCalled(); + expect(networking.tagsSpendLogsCall).not.toHaveBeenCalled(); + expect(networking.allTagNamesCall).not.toHaveBeenCalled(); + expect(networking.adminspendByProvider).not.toHaveBeenCalled(); + expect(networking.adminGlobalActivity).not.toHaveBeenCalled(); + expect(networking.adminGlobalActivityPerModel).not.toHaveBeenCalled(); + }); + }, + ); + + describe("the admin-only gate", () => { + it("proves the flush window is wide enough to catch a leaked request", async () => { + renderUsage({ userRole: "Admin" }); + + await flushPendingRequests(); + + expect(networking.getProxyUISettings).toHaveBeenCalled(); + expect(networking.adminSpendLogsCall).toHaveBeenCalled(); + expect(networking.tagsSpendLogsCall).toHaveBeenCalled(); + expect(networking.adminGlobalActivity).toHaveBeenCalled(); + }); + + it("still lets an admin through, so the notice is a real gate and not a dead branch", async () => { + renderUsage({ userRole: "Admin" }); expect(await screen.findByRole("tab", { name: "All Up" })).toBeInTheDocument(); - expect(screen.queryByRole("tab", { name: "Team Based Usage" })).not.toBeInTheDocument(); - expect(screen.queryByRole("tab", { name: "Customer Usage" })).not.toBeInTheDocument(); - expect(screen.queryByRole("tab", { name: "Tag Based Usage" })).not.toBeInTheDocument(); - + expect(screen.queryByText(/Proxy-wide usage is only available to admin users/i)).not.toBeInTheDocument(); await waitFor(() => expect(networking.adminSpendLogsCall).toHaveBeenCalled()); - expect(networking.teamSpendLogsCall).not.toHaveBeenCalled(); - expect(networking.adminTopEndUsersCall).not.toHaveBeenCalled(); + }); + + it("does not put the session token in the provider spend query", async () => { + renderUsage({ userRole: "Admin", token: "session-jwt-value" }); + + await waitFor(() => expect(networking.adminspendByProvider).toHaveBeenCalled()); + const callArgs = networking.adminspendByProvider.mock.calls[0]; + expect(callArgs).not.toContain("session-jwt-value"); + expect(callArgs[0]).toBe("sk-test"); }); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.tsx index 3d55f9bb698..5b2f8547822 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.tsx @@ -37,6 +37,7 @@ import { } from "@/components/networking"; import TopKeyView from "@/components/UsagePage/components/EntityUsage/TopKeyView"; import { MoneyCell } from "@/components/shared/table_cells"; +import { hasCapability } from "@/utils/capabilities"; import { formatNumberWithCommas } from "@/utils/dataUtils"; interface UsagePageProps { @@ -90,6 +91,7 @@ const TeamSpendBarList: React.FC<{ data: TeamSpendTotal[] }> = ({ data }) => { }; const UsagePage: React.FC = ({ accessToken, token, userRole, userID, keys, premiumUser }) => { + const canViewGlobalSpend = hasCapability(userRole, "viewGlobalSpend"); const currentDate = new Date(); const [keySpendData, setKeySpendData] = useState([]); const [topKeys, setTopKeys] = useState([]); @@ -155,8 +157,11 @@ const UsagePage: React.FC = ({ accessToken, token, userRole, use }; useEffect(() => { + if (!canViewGlobalSpend) { + return; + } updateTagSpendData(dateValue.from, dateValue.to); - }, [dateValue, selectedTags]); + }, [canViewGlobalSpend, dateValue, selectedTags]); const updateEndUserData = async ( startTime: Date | undefined, @@ -319,10 +324,7 @@ const UsagePage: React.FC = ({ accessToken, token, userRole, use const fetchProviderSpend = () => fetchAndSetData( - () => - accessToken && token - ? adminspendByProvider(accessToken, token, startTime, endTime) - : Promise.reject("No access token or token"), + () => (accessToken ? adminspendByProvider(accessToken, startTime, endTime) : Promise.reject("No access token")), setSpendByProvider, "Error fetching provider spend", ); @@ -467,6 +469,9 @@ const UsagePage: React.FC = ({ accessToken, token, userRole, use useEffect(() => { const initlizeUsageData = async () => { + if (!canViewGlobalSpend) { + return; + } if (accessToken && token && userRole && userID) { const proxy_settings: ProxySettings | undefined = await fetchProxySettings(); if (proxy_settings) { @@ -493,7 +498,24 @@ const UsagePage: React.FC = ({ accessToken, token, userRole, use }; initlizeUsageData(); - }, [accessToken, token, userRole, userID, startTime, endTime]); + }, [canViewGlobalSpend, accessToken, token, userRole, userID, startTime, endTime]); + + if (!canViewGlobalSpend) { + return ( +
+ + + Usage + + +

+ Proxy-wide usage is only available to admin users. Your own usage is on the Usage page. +

+
+
+
+ ); + } if (proxySettings?.DISABLE_EXPENSIVE_DB_QUERIES) { return ( diff --git a/ui/litellm-dashboard/src/components/key_team_helpers/filter_helpers.test.ts b/ui/litellm-dashboard/src/components/key_team_helpers/filter_helpers.test.ts index 637325ad98d..15c45153026 100644 --- a/ui/litellm-dashboard/src/components/key_team_helpers/filter_helpers.test.ts +++ b/ui/litellm-dashboard/src/components/key_team_helpers/filter_helpers.test.ts @@ -1,11 +1,12 @@ -import { describe, expect, it, vi } from "vitest"; -import { fetchTeamFilterOptions } from "./filter_helpers"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { fetchAllTeams, fetchTeamFilterOptions } from "./filter_helpers"; const mockKeyListCall = vi.fn(); +const mockTeamListCall = vi.fn(); vi.mock("@/components/networking", () => ({ keyListCall: (...args: unknown[]) => mockKeyListCall(...args), - teamListCall: vi.fn(), + teamListCall: (...args: unknown[]) => mockTeamListCall(...args), organizationListCall: vi.fn(), })); @@ -78,3 +79,39 @@ describe("fetchTeamFilterOptions", () => { expect(result).toEqual({ keyAliases: [], organizationIds: [], userIds: [] }); }); }); + +describe("fetchAllTeams", () => { + beforeEach(() => { + mockTeamListCall.mockReset(); + }); + + it("forwards the scoping user id to /team/list and returns the rows it answers with", async () => { + mockTeamListCall.mockResolvedValue([{ team_id: "team-a" }, { team_id: "team-b" }]); + + const teams = await fetchAllTeams("tok-123", null, "member-7"); + + expect(mockTeamListCall).toHaveBeenCalledWith("tok-123", null, "member-7"); + expect(teams.map((team) => team.team_id)).toEqual(["team-a", "team-b"]); + }); + + it("sends no user id when the caller is entitled to the broad list", async () => { + mockTeamListCall.mockResolvedValue([]); + + await fetchAllTeams("tok-123"); + + expect(mockTeamListCall).toHaveBeenCalledWith("tok-123", null, null); + }); + + it("keeps the organization filter independent of the scoping user id", async () => { + mockTeamListCall.mockResolvedValue([]); + + await fetchAllTeams("tok-123", "org-1", "member-7"); + + expect(mockTeamListCall).toHaveBeenCalledWith("tok-123", "org-1", "member-7"); + }); + + it("returns an empty list without calling the endpoint when there is no access token", async () => { + expect(await fetchAllTeams(null, null, "member-7")).toEqual([]); + expect(mockTeamListCall).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/key_team_helpers/filter_helpers.ts b/ui/litellm-dashboard/src/components/key_team_helpers/filter_helpers.ts index fb701b4656b..7eef4d3a8b3 100644 --- a/ui/litellm-dashboard/src/components/key_team_helpers/filter_helpers.ts +++ b/ui/litellm-dashboard/src/components/key_team_helpers/filter_helpers.ts @@ -114,9 +114,15 @@ export const fetchTeamFilterOptions = async ( * Fetches all teams across all pages * @param accessToken The access token for API authentication * @param organizationId Optional organization ID to filter teams + * @param userID Scopes the list to that user's teams. Required for roles the endpoint + * does not grant a broad list to; see `teamListScopeUserId` * @returns Array of all teams */ -export const fetchAllTeams = async (accessToken: string | null, organizationId?: string | null): Promise => { +export const fetchAllTeams = async ( + accessToken: string | null, + organizationId?: string | null, + userID?: string | null, +): Promise => { if (!accessToken) return []; try { @@ -125,7 +131,7 @@ export const fetchAllTeams = async (accessToken: string | null, organizationId?: let hasMorePages = true; while (hasMorePages) { - const response = await teamListCall(accessToken, organizationId || null, null); + const response = await teamListCall(accessToken, organizationId || null, userID ?? null); // Add teams from this page allTeams = [...allTeams, ...response]; diff --git a/ui/litellm-dashboard/src/components/leftnav.test.tsx b/ui/litellm-dashboard/src/components/leftnav.test.tsx index 3c2047b8d75..f074624dcf2 100644 --- a/ui/litellm-dashboard/src/components/leftnav.test.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.test.tsx @@ -8,6 +8,7 @@ vi.mock("../utils/roles", async (importOriginal) => { return { ...actual, all_admin_roles: ["admin", "admin_viewer"], + old_admin_roles: ["admin", "admin_viewer"], internalUserRoles: ["internal"], rolesWithWriteAccess: ["admin", "internal"], rolesAllowedToViewWriteScopedPages: ["admin", "internal", "admin_viewer"], @@ -273,6 +274,30 @@ describe("Sidebar (leftnav)", () => { }); expect(screen.queryByText("Prompts")).not.toBeInTheDocument(); }); + + it("should hide Old Usage from internal users while keeping other Experimental children", async () => { + mockUseAuthorized.mockReturnValue(internalAuth); + renderWithProviders(); + + act(() => { + fireEvent.click(screen.getByText("Experimental")); + }); + await waitFor(() => { + expect(screen.getByText("API Playground")).toBeInTheDocument(); + }); + expect(screen.queryByText("Old Usage")).not.toBeInTheDocument(); + }); + + it("should show Old Usage to admins", async () => { + renderWithProviders(); + + act(() => { + fireEvent.click(screen.getByText("Experimental")); + }); + await waitFor(() => { + expect(screen.getByText("Old Usage")).toBeInTheDocument(); + }); + }); }); it("should show Organizations tab for organization admins", () => { diff --git a/ui/litellm-dashboard/src/components/leftnav.tsx b/ui/litellm-dashboard/src/components/leftnav.tsx index 2e199db028d..42f427665d4 100644 --- a/ui/litellm-dashboard/src/components/leftnav.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.tsx @@ -288,7 +288,13 @@ const menuGroups: MenuGroup[] = [ icon: , roles: all_admin_roles, }, - { key: "4", page: "usage", label: "Old Usage", icon: }, + { + key: "4", + page: "usage", + label: "Old Usage", + icon: , + roles: rolesWithCapability("viewGlobalSpend"), + }, ], }, ], diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 4a396fa83fc..321dfd6b2a6 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -2099,7 +2099,6 @@ export const adminTopEndUsersCall = async ( export const adminspendByProvider = async ( accessToken: string, - keyToken: string | null, startTime: string | undefined, endTime: string | undefined, ) => { @@ -2108,7 +2107,6 @@ export const adminspendByProvider = async ( accessToken, query: { ...(startTime && endTime ? { start_date: startTime, end_date: endTime } : {}), - ...(keyToken ? { api_key: keyToken } : {}), }, }); return data; diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test-d.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test-d.tsx new file mode 100644 index 00000000000..7bbd4f918cd --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test-d.tsx @@ -0,0 +1,67 @@ +import type { ColumnDef, PaginationState, RowSelectionState, SortingState } from "@tanstack/react-table"; + +import { DataTable } from "./DataTable"; + +interface Row { + id: string; + name: string; +} + +const data: Row[] = []; +const columns: ColumnDef[] = []; +const sorting: SortingState = [{ id: "name", desc: false }]; +const pagination: PaginationState = { pageIndex: 0, pageSize: 10 }; +const rowSelection: RowSelectionState = { r1: true }; +const noop = () => {}; + +export const uncontrolled = ; + +export const controlled = ( + +); + +export const clientSortingWithServerPagination = ( + +); + +// @ts-expect-error sortingMode="server" requires `sorting` and `onSortingChange` +export const serverSortingWithoutState = ; + +// @ts-expect-error paginationMode="server" requires `pagination`, `onPaginationChange` and `rowCount` +export const serverPaginationWithoutState = ; + +// @ts-expect-error filterMode="server" requires `columnFilters` and `onColumnFiltersChange` +export const serverFilteringWithoutState = ; + +export const bothSortingSources = ( + // @ts-expect-error `defaultSorting` seeds uncontrolled sorting, so it cannot pair with a controlled `sorting` + +); + +export const selectionWithoutHandler = ( + // @ts-expect-error a controlled `rowSelection` needs `onRowSelectionChange` or selection changes are dropped + +); diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx index 8555a10c326..3afdd2849ae 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.test.tsx @@ -337,14 +337,6 @@ describe("DataTable filtering", () => { ); expect(names()).toEqual(["Charlie", "Alice", "Bob"]); }); - - it("throws when server filtering is missing required props", () => { - const spy = vi.spyOn(console, "error").mockImplementation(() => {}); - expect(() => render()).toThrow( - /filterMode='server'/, - ); - spy.mockRestore(); - }); }); describe("DataTable loading", () => { @@ -664,36 +656,3 @@ describe("DataTable layout", () => { expect(container.querySelector("thead")?.className).not.toContain("bg-background"); }); }); - -describe("DataTable misconfiguration guards", () => { - it("throws when server sorting is missing required props", () => { - const spy = vi.spyOn(console, "error").mockImplementation(() => {}); - expect(() => render()).toThrow( - /sortingMode='server'/, - ); - spy.mockRestore(); - }); - - it("throws when server pagination is missing required props", () => { - const spy = vi.spyOn(console, "error").mockImplementation(() => {}); - expect(() => render()).toThrow( - /paginationMode='server'/, - ); - spy.mockRestore(); - }); - - it("throws when both defaultSorting and sorting are provided", () => { - const spy = vi.spyOn(console, "error").mockImplementation(() => {}); - expect(() => - render( - , - ), - ).toThrow(/defaultSorting/); - spy.mockRestore(); - }); -}); diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx index 8cd0e25dfc4..2f47887d01f 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTable.tsx @@ -42,7 +42,15 @@ import { cn } from "@/lib/cva.config"; import "./columnMeta"; import { DataTablePagination, DEFAULT_PAGE_SIZE_OPTIONS } from "./DataTablePagination"; -import type { ColumnPinnedSide, DataTableProps, DataTableSize, FilterMode, PaginationMode, SortingMode } from "./types"; +import type { + ColumnPinnedSide, + DataTableProps, + DataTableResolvedProps, + DataTableSize, + FilterMode, + PaginationMode, + SortingMode, +} from "./types"; const INTERACTIVE_SELECTOR = "button, a, input, select, textarea, [role=checkbox], [data-row-click-exempt]"; @@ -64,47 +72,6 @@ const FILL_CLASSES = { const NO_FILL_CLASSES = { outer: "", frame: "", body: "", header: "" } as const; -export class DataTableConfigError extends Error { - constructor(messages: readonly string[]) { - super(`DataTable misconfiguration:\n- ${messages.join("\n- ")}`); - this.name = "DataTableConfigError"; - } -} - -export function validateDataTableConfig( - props: DataTableProps, -): readonly string[] { - const serverSortingIncomplete = - props.sortingMode === "server" && (props.sorting === undefined || props.onSortingChange === undefined); - - const serverPaginationPropsMissing = - props.pagination === undefined || props.onPaginationChange === undefined || props.rowCount === undefined; - const serverPaginationIncomplete = props.paginationMode === "server" && serverPaginationPropsMissing; - - const serverFilteringIncomplete = - props.filterMode === "server" && (props.columnFilters === undefined || props.onColumnFiltersChange === undefined); - - const bothSortingSources = props.defaultSorting !== undefined && props.sorting !== undefined; - const bothFilterSources = props.defaultColumnFilters !== undefined && props.columnFilters !== undefined; - - const controlledSelectionIncomplete = props.rowSelection !== undefined && props.onRowSelectionChange === undefined; - - return [ - serverSortingIncomplete ? "sortingMode='server' requires both `sorting` and `onSortingChange`." : null, - serverPaginationIncomplete - ? "paginationMode='server' requires `pagination`, `onPaginationChange`, and `rowCount`." - : null, - serverFilteringIncomplete ? "filterMode='server' requires both `columnFilters` and `onColumnFiltersChange`." : null, - bothSortingSources ? "Provide either `defaultSorting` (uncontrolled) or `sorting` (controlled), not both." : null, - bothFilterSources - ? "Provide either `defaultColumnFilters` (uncontrolled) or `columnFilters` (controlled), not both." - : null, - controlledSelectionIncomplete - ? "Controlled `rowSelection` requires `onRowSelectionChange`; without it selection changes are dropped." - : null, - ].filter((message): message is string => message !== null); -} - function columnDefId(column: ColumnDef): string | undefined { if ("id" in column && typeof column.id === "string") { return column.id; @@ -442,7 +409,9 @@ function useControllable( return { value: internal, onChange: setInternal }; } -function useDataTableInstance(props: DataTableProps): Table { +function useDataTableInstance( + props: DataTableResolvedProps, +): Table { const { data, columns, @@ -532,14 +501,7 @@ function useDataTableInstance(props: DataTablePro } export function DataTable(props: DataTableProps) { - // Validate once at construction so a misconfig surfaces immediately instead of on every render. - useState(() => { - const errors = validateDataTableConfig(props); - if (errors.length > 0) { - throw new DataTableConfigError(errors); - } - return null; - }); + const resolved: DataTableResolvedProps = props; const { isLoading = false, @@ -559,9 +521,9 @@ export function DataTable(props: DataTableProps { await user.click(rowBox("m1")); expect(selectedCount()).toHaveTextContent("1"); }); - - it("rejects controlled rowSelection without onRowSelectionChange", () => { - const errors = validateDataTableConfig({ data, columns, rowSelection: { m1: true } }); - - expect(errors).toContain( - "Controlled `rowSelection` requires `onRowSelectionChange`; without it selection changes are dropped.", - ); - }); - - it("does not complain when selection is left uncontrolled", () => { - expect(validateDataTableConfig({ data, columns })).toHaveLength(0); - }); }); diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/index.ts b/ui/litellm-dashboard/src/components/shared/DataTable/index.ts index 62ddd1b0742..39a887ba948 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/index.ts +++ b/ui/litellm-dashboard/src/components/shared/DataTable/index.ts @@ -1,6 +1,6 @@ import "./columnMeta"; -export { DataTable, DataTableConfigError, validateDataTableConfig } from "./DataTable"; +export { DataTable } from "./DataTable"; export { DataTableFilterDrawer, DataTableFilterField, type FilterDraft } from "./DataTableFilterDrawer"; export { DataTablePagination, DEFAULT_PAGE_SIZE_OPTIONS } from "./DataTablePagination"; export { createSelectionColumn } from "./DataTableSelectionColumn"; @@ -17,6 +17,7 @@ export type { ColumnPinnedSide, ColumnResizeMode, DataTableProps, + DataTableResolvedProps, DataTableSize, FilterMode, PaginationMode, diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/types.ts b/ui/litellm-dashboard/src/components/shared/DataTable/types.ts index dd578f4df45..c767a0a64c0 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/types.ts +++ b/ui/litellm-dashboard/src/components/shared/DataTable/types.ts @@ -21,7 +21,7 @@ export type DataTableSize = "compact" | "default"; export type ColumnPinnedSide = "left" | "right"; export type DataTableSkeletonShape = "text" | "twoLine" | "badge" | "chips" | "meter"; -export interface DataTableProps { +export interface DataTableResolvedProps { data: TData[]; columns: ColumnDef[]; getRowId?: (row: TData, index: number, parent?: Row) => string; @@ -81,3 +81,73 @@ export interface DataTableProps { paginationSlot?: (table: Table) => React.ReactNode; footer?: (table: Table) => React.ReactNode; } + +type DataTableBaseProps = Omit< + DataTableResolvedProps, + | "sortingMode" + | "sorting" + | "onSortingChange" + | "defaultSorting" + | "paginationMode" + | "pagination" + | "onPaginationChange" + | "rowCount" + | "filterMode" + | "columnFilters" + | "onColumnFiltersChange" + | "defaultColumnFilters" + | "rowSelection" + | "onRowSelectionChange" +>; + +type SortingProps = + | { + sorting: SortingState; + onSortingChange: OnChangeFn; + sortingMode?: SortingMode; + defaultSorting?: never; + } + | { + sortingMode?: Exclude; + sorting?: never; + onSortingChange?: never; + defaultSorting?: SortingState; + }; + +type PaginationProps = + | { + paginationMode: "server"; + pagination: PaginationState; + onPaginationChange: OnChangeFn; + rowCount: number; + } + | { + paginationMode?: Exclude; + pagination?: PaginationState; + onPaginationChange?: OnChangeFn; + rowCount?: number; + }; + +type FilterProps = + | { + columnFilters: ColumnFiltersState; + onColumnFiltersChange: OnChangeFn; + filterMode?: FilterMode; + defaultColumnFilters?: never; + } + | { + filterMode?: Exclude; + columnFilters?: never; + onColumnFiltersChange?: never; + defaultColumnFilters?: ColumnFiltersState; + }; + +type RowSelectionProps = + | { rowSelection: RowSelectionState; onRowSelectionChange: OnChangeFn } + | { rowSelection?: never; onRowSelectionChange?: OnChangeFn }; + +export type DataTableProps = DataTableBaseProps & + SortingProps & + PaginationProps & + FilterProps & + RowSelectionProps; diff --git a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx index 080bf6b380a..17d26dc00f3 100644 --- a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.test.tsx @@ -26,6 +26,8 @@ vi.mock("@/components/key_team_helpers/filter_helpers", () => ({ })); import { uiSpendLogsCall } from "../networking"; +import { fetchAllTeams } from "@/components/key_team_helpers/filter_helpers"; +import type { Team } from "../key_team_helpers/key_list"; const emptyResponse: PaginatedResponse = { data: [], @@ -198,6 +200,39 @@ describe("useLogFilterLogic", () => { }); }); + describe("team filter list scope", () => { + const callerTeams = [{ team_id: "team-a" }, { team_id: "team-b" }] as Team[]; + + it("scopes /team/list to an internal user and still surfaces their teams", async () => { + vi.mocked(fetchAllTeams).mockResolvedValue(callerTeams); + + const { result } = renderFilterHook({ userRole: "Internal User", userID: "member-7" }); + + await waitFor(() => expect(fetchAllTeams).toHaveBeenCalled()); + expect(fetchAllTeams).toHaveBeenCalledWith("test-token", null, "member-7"); + await waitFor(() => expect(result.current.allTeams).toEqual(callerTeams)); + }); + + it("scopes /team/list for an internal viewer", async () => { + vi.mocked(fetchAllTeams).mockResolvedValue(callerTeams); + + renderFilterHook({ userRole: "Internal Viewer", userID: "member-7" }); + + await waitFor(() => expect(fetchAllTeams).toHaveBeenCalledWith("test-token", null, "member-7")); + }); + + it.each(["Admin", "Admin Viewer", "Org Admin"])( + "leaves /team/list unscoped for %s so the broad list survives", + async (userRole) => { + vi.mocked(fetchAllTeams).mockResolvedValue(callerTeams); + + renderFilterHook({ userRole, userID: "member-7" }); + + await waitFor(() => expect(fetchAllTeams).toHaveBeenCalledWith("test-token", null, null)); + }, + ); + }); + it("returns an empty payload and does not crash when the call fails", async () => { vi.mocked(uiSpendLogsCall).mockRejectedValue(new Error("boom")); const { result } = renderFilterHook(); diff --git a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx index 8244066cadb..e1089c6a16c 100644 --- a/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/log_filter_logic.tsx @@ -4,6 +4,7 @@ import type { ColumnFiltersState, PaginationState, SortingState } from "@tanstac import { uiSpendLogsCall } from "../networking"; import { Team } from "../key_team_helpers/key_list"; import { fetchAllTeams } from "../../components/key_team_helpers/filter_helpers"; +import { teamListScopeUserId } from "../../utils/roles"; import { defaultPageSize } from "../constants"; import { LOGS_SORT_FIELD_MAP, type LogEntry, type LogsSortField } from "./columns"; @@ -194,11 +195,13 @@ export function useLogFilterLogic({ total_pages: 0, }; + const teamListUserID = teamListScopeUserId(userRole, userID); + const allTeamsQueryOptions: UseQueryOptions = { - queryKey: ["allTeamsForLogFilters", accessToken], + queryKey: ["allTeamsForLogFilters", accessToken, teamListUserID], queryFn: async () => { if (!accessToken) return []; - const teamsData = await fetchAllTeams(accessToken); + const teamsData = await fetchAllTeams(accessToken, null, teamListUserID); return teamsData || []; }, enabled: !!accessToken, diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index fa8731d7a16..6323516b126 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -26758,6 +26758,11 @@ export interface components { } | null; /** Adaptive Router Default Model */ adaptive_router_default_model?: string | null; + /** + * Allow Client Keepalive Override + * @default false + */ + allow_client_keepalive_override: boolean | null; /** Annotation Cost Per Page */ annotation_cost_per_page?: number | null; /** Api Base */ @@ -26906,6 +26911,8 @@ export interface components { input_cost_per_video_token?: number | null; /** Itpm */ itpm?: number | null; + /** Keepalive Seconds */ + keepalive_seconds?: number | null; /** Litellm Credential Name */ litellm_credential_name?: string | null; /** Litellm Trace Id */ @@ -35429,6 +35436,11 @@ export interface components { } | null; /** Adaptive Router Default Model */ adaptive_router_default_model?: string | null; + /** + * Allow Client Keepalive Override + * @default false + */ + allow_client_keepalive_override: boolean | null; /** Annotation Cost Per Page */ annotation_cost_per_page?: number | null; /** Api Base */ @@ -35577,6 +35589,8 @@ export interface components { input_cost_per_video_token?: number | null; /** Itpm */ itpm?: number | null; + /** Keepalive Seconds */ + keepalive_seconds?: number | null; /** Litellm Credential Name */ litellm_credential_name?: string | null; /** Litellm Trace Id */ diff --git a/ui/litellm-dashboard/src/utils/capabilities.test.ts b/ui/litellm-dashboard/src/utils/capabilities.test.ts index 5a248c09200..d53b450c730 100644 --- a/ui/litellm-dashboard/src/utils/capabilities.test.ts +++ b/ui/litellm-dashboard/src/utils/capabilities.test.ts @@ -1,6 +1,7 @@ import { describe, expect, it } from "vitest"; import { hasCapability, rolesWithCapability, type Capability } from "./capabilities"; +import { effectiveSessionRole } from "./roles"; const ADMIN_ROLES = ["Admin", "Admin Viewer", "proxy_admin", "proxy_admin_viewer"]; const NON_ADMIN_ROLES = [ @@ -67,6 +68,39 @@ describe("hasCapability for organization admins", () => { ); expect(orgAdminCapabilities).toEqual(["viewDeletedTeams"]); }); + + it("does not let the org-admin allowance reopen the proxy-admin-only viewGlobalSpend gate", () => { + expect(hasCapability(SESSION_ROLE_AN_ORG_ADMIN_ACTUALLY_CARRIES, "viewGlobalSpend", true)).toBe(false); + expect(hasCapability("Org Admin", "viewGlobalSpend", true)).toBe(false); + }); +}); + +describe("hasCapability - viewGlobalSpend", () => { + it.each(ADMIN_ROLES)("should grant it to %s", (role) => { + expect(hasCapability(role, "viewGlobalSpend")).toBe(true); + }); + + it.each([...NON_ADMIN_ROLES, "internal_user_viewer", "org_admin"])("should deny it to %s", (role) => { + expect(hasCapability(role, "viewGlobalSpend")).toBe(false); + }); + + it("should deny it to every role an org admin or team admin can present at runtime", () => { + const orgAdminSessionRole = effectiveSessionRole("internal_user"); + const teamAdminSessionRole = effectiveSessionRole("internal_user"); + + expect(orgAdminSessionRole).toBe("Internal User"); + expect(hasCapability(orgAdminSessionRole, "viewGlobalSpend")).toBe(false); + expect(hasCapability(teamAdminSessionRole, "viewGlobalSpend")).toBe(false); + }); + + it.each([ + ["proxy_admin", true], + ["proxy_admin_viewer", true], + ["internal_user", false], + ["internal_user_viewer", false], + ] as const)("should match the backend for a %s session", (rawRole, expected) => { + expect(hasCapability(effectiveSessionRole(rawRole), "viewGlobalSpend")).toBe(expected); + }); }); describe("rolesWithCapability", () => { diff --git a/ui/litellm-dashboard/src/utils/capabilities.ts b/ui/litellm-dashboard/src/utils/capabilities.ts index c529c0c43ef..ebd8b8f06a3 100644 --- a/ui/litellm-dashboard/src/utils/capabilities.ts +++ b/ui/litellm-dashboard/src/utils/capabilities.ts @@ -1,4 +1,6 @@ -import { all_admin_roles } from "./roles"; +import { all_admin_roles, old_admin_roles } from "./roles"; + +const proxyAdminOnlyRoles = [...old_admin_roles, "proxy_admin", "proxy_admin_viewer"]; const CAPABILITY_ROLES = { viewToolPolicies: all_admin_roles, @@ -8,6 +10,7 @@ const CAPABILITY_ROLES = { viewPrompts: all_admin_roles, viewOrganizationUsage: all_admin_roles, viewAgentUsage: all_admin_roles, + viewGlobalSpend: proxyAdminOnlyRoles, } as const satisfies Record; export type Capability = keyof typeof CAPABILITY_ROLES; diff --git a/ui/litellm-dashboard/src/utils/roles.test.ts b/ui/litellm-dashboard/src/utils/roles.test.ts index 3fe3a1a0f0b..b2d7c2c75e5 100644 --- a/ui/litellm-dashboard/src/utils/roles.test.ts +++ b/ui/litellm-dashboard/src/utils/roles.test.ts @@ -1,5 +1,6 @@ import { describe, it, expect } from "vitest"; import { + all_admin_roles, effectiveSessionRole, isAdminRole, isOrgAdminForAnyOrg, @@ -10,6 +11,7 @@ import { isViewOnlySessionRole, rolesAllowedToViewWriteScopedPages, rolesWithWriteAccess, + teamListScopeUserId, } from "./roles"; import { Organization, Team } from "@/components/networking"; @@ -290,4 +292,37 @@ describe("roles", () => { expect(isViewOnlySessionRole("proxy_admin_viewer")).toBe(true); }); }); + + describe("teamListScopeUserId", () => { + const SESSION_USER_ID = "user-1"; + + it.each(["proxy_admin", "proxy_admin_viewer", "org_admin"])( + "leaves %s unscoped so the endpoint keeps returning its broad list", + (rawRole) => { + expect(teamListScopeUserId(effectiveSessionRole(rawRole), SESSION_USER_ID)).toBeNull(); + }, + ); + + it.each(["internal_user", "internal_user_viewer", "internal_viewer", "app_user"])( + "scopes %s to its own user id, which is what the endpoint authorizes on", + (rawRole) => { + expect(teamListScopeUserId(effectiveSessionRole(rawRole), SESSION_USER_ID)).toBe(SESSION_USER_ID); + }, + ); + + it("also accepts the Admin Viewer label that formatUserRole emits", () => { + expect(teamListScopeUserId("Admin Viewer", SESSION_USER_ID)).toBeNull(); + }); + + it("scopes an unknown or absent role rather than assuming a broad list", () => { + expect(teamListScopeUserId(null, SESSION_USER_ID)).toBe(SESSION_USER_ID); + expect(teamListScopeUserId("Undefined Role", SESSION_USER_ID)).toBe(SESSION_USER_ID); + }); + + it("keeps Org Admin broad even though all_admin_roles carries only the raw org_admin", () => { + expect(all_admin_roles).not.toContain(effectiveSessionRole("org_admin")); + expect(isAdminRole(effectiveSessionRole("org_admin"))).toBe(false); + expect(teamListScopeUserId(effectiveSessionRole("org_admin"), SESSION_USER_ID)).toBeNull(); + }); + }); }); diff --git a/ui/litellm-dashboard/src/utils/roles.ts b/ui/litellm-dashboard/src/utils/roles.ts index 36e7bf3beb1..b6d84b9365d 100644 --- a/ui/litellm-dashboard/src/utils/roles.ts +++ b/ui/litellm-dashboard/src/utils/roles.ts @@ -100,3 +100,12 @@ export const effectiveSessionRole = (rawUserRole?: string): string => { export const isViewOnlySessionRole = (rawUserRole?: string): boolean => viewOnlyRawRoles.includes(rawUserRole?.toLowerCase() ?? ""); + +// Session roles (the value `useAuthorized().userRole` supplies) that /team/list and +// /v2/team/list already answer with a broad list: proxy-wide for admins, org-wide for +// org admins. Sending a user_id for those narrows the response to direct memberships, +// so only the roles the endpoints would otherwise reject carry one. +const sessionRolesWithBroadTeamList: string[] = ["Admin", "Admin Viewer", "Org Admin"]; + +export const teamListScopeUserId = (userRole: string | null, userId: string | null): string | null => + sessionRolesWithBroadTeamList.includes(userRole ?? "") ? null : userId; diff --git a/ui/litellm-dashboard/vitest.config.ts b/ui/litellm-dashboard/vitest.config.ts index 19af41ab8ed..469f7fa3520 100644 --- a/ui/litellm-dashboard/vitest.config.ts +++ b/ui/litellm-dashboard/vitest.config.ts @@ -29,6 +29,7 @@ const config: ViteUserConfig = { exclude: [ "**/*.d.ts", "**/*.test.*", + "**/*.test-d.*", "**/*.spec.*", "tests/**", @@ -45,6 +46,10 @@ const config: ViteUserConfig = { }, exclude: ["node_modules/**"], include: ["src/**/*.test.ts", "src/**/*.test.tsx", "tests/**/*.test.ts", "tests/**/*.test.tsx"], + typecheck: { + include: ["src/**/*.test-d.ts", "src/**/*.test-d.tsx"], + ignoreSourceErrors: true, + }, }, resolve: { alias: {