Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/vibrant-booth-d4258b

# Conflicts:
#	ui/litellm-dashboard/src/utils/capabilities.test.ts
This commit is contained in:
Yuneng Jiang 2026-08-10 16:50:50 -07:00
commit d46ef9aeb4
No known key found for this signature in database
36 changed files with 1957 additions and 163 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -3399,6 +3399,8 @@ all_litellm_params = (
+ [
"metadata",
"litellm_metadata",
"keepalive_seconds",
"allow_client_keepalive_override",
"litellm_trace_id",
"litellm_request_debug",
"guardrails",

View file

@ -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<version>`, 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<provider_version>`
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<provider_version>`. 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<MAJOR>.<MINOR>.<PATCH>` (e.g., `v0.1.2`, `v1.0.0`)

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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<Team[]> => {
const ALL_TEAMS_PAGE_SIZE = 100;
const fetchAllTeamsPaged = async (accessToken: string): Promise<Team[]> => {
const firstPage: TeamsResponse = await teamListCall(accessToken, 1, ALL_TEAMS_PAGE_SIZE);
const fetchAllTeamsPaged = async (accessToken: string, userID: string | null): Promise<Team[]> => {
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<Team[]> => {
const { accessToken } = useAuthorized();
const { accessToken, userId, userRole } = useAuthorized();
const scopedUserID = teamListScopeUserId(userRole, userId);
return useQuery<Team[]>({
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,
});

View file

@ -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<React.ComponentProps<typeof UsagePage>>
/>,
);
// 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");
});
});
});

View file

@ -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<UsagePageProps> = ({ accessToken, token, userRole, userID, keys, premiumUser }) => {
const canViewGlobalSpend = hasCapability(userRole, "viewGlobalSpend");
const currentDate = new Date();
const [keySpendData, setKeySpendData] = useState<any[]>([]);
const [topKeys, setTopKeys] = useState<any[]>([]);
@ -155,8 +157,11 @@ const UsagePage: React.FC<UsagePageProps> = ({ 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<UsagePageProps> = ({ 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<UsagePageProps> = ({ 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<UsagePageProps> = ({ accessToken, token, userRole, use
};
initlizeUsageData();
}, [accessToken, token, userRole, userID, startTime, endTime]);
}, [canViewGlobalSpend, accessToken, token, userRole, userID, startTime, endTime]);
if (!canViewGlobalSpend) {
return (
<div className="w-full p-8">
<Card>
<CardHeader>
<CardTitle>Usage</CardTitle>
</CardHeader>
<CardContent>
<p className="text-sm text-muted-foreground">
Proxy-wide usage is only available to admin users. Your own usage is on the Usage page.
</p>
</CardContent>
</Card>
</div>
);
}
if (proxySettings?.DISABLE_EXPENSIVE_DB_QUERIES) {
return (

View file

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

View file

@ -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<Team[]> => {
export const fetchAllTeams = async (
accessToken: string | null,
organizationId?: string | null,
userID?: string | null,
): Promise<Team[]> => {
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];

View file

@ -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(<Sidebar {...defaultProps} />);
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(<Sidebar {...defaultProps} />);
act(() => {
fireEvent.click(screen.getByText("Experimental"));
});
await waitFor(() => {
expect(screen.getByText("Old Usage")).toBeInTheDocument();
});
});
});
it("should show Organizations tab for organization admins", () => {

View file

@ -288,7 +288,13 @@ const menuGroups: MenuGroup[] = [
icon: <Tags {...ICON} />,
roles: all_admin_roles,
},
{ key: "4", page: "usage", label: "Old Usage", icon: <BarChart3 {...ICON} /> },
{
key: "4",
page: "usage",
label: "Old Usage",
icon: <BarChart3 {...ICON} />,
roles: rolesWithCapability("viewGlobalSpend"),
},
],
},
],

View file

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

View file

@ -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<Row, unknown>[] = [];
const sorting: SortingState = [{ id: "name", desc: false }];
const pagination: PaginationState = { pageIndex: 0, pageSize: 10 };
const rowSelection: RowSelectionState = { r1: true };
const noop = () => {};
export const uncontrolled = <DataTable data={data} columns={columns} defaultSorting={sorting} />;
export const controlled = (
<DataTable
data={data}
columns={columns}
sortingMode="server"
sorting={sorting}
onSortingChange={noop}
paginationMode="server"
pagination={pagination}
onPaginationChange={noop}
rowCount={0}
filterMode="server"
columnFilters={[]}
onColumnFiltersChange={noop}
rowSelection={rowSelection}
onRowSelectionChange={noop}
/>
);
export const clientSortingWithServerPagination = (
<DataTable
data={data}
columns={columns}
defaultSorting={sorting}
paginationMode="server"
pagination={pagination}
onPaginationChange={noop}
rowCount={0}
/>
);
// @ts-expect-error sortingMode="server" requires `sorting` and `onSortingChange`
export const serverSortingWithoutState = <DataTable data={data} columns={columns} sortingMode="server" />;
// @ts-expect-error paginationMode="server" requires `pagination`, `onPaginationChange` and `rowCount`
export const serverPaginationWithoutState = <DataTable data={data} columns={columns} paginationMode="server" />;
// @ts-expect-error filterMode="server" requires `columnFilters` and `onColumnFiltersChange`
export const serverFilteringWithoutState = <DataTable data={data} columns={columns} filterMode="server" />;
export const bothSortingSources = (
// @ts-expect-error `defaultSorting` seeds uncontrolled sorting, so it cannot pair with a controlled `sorting`
<DataTable data={data} columns={columns} defaultSorting={sorting} sorting={sorting} onSortingChange={noop} />
);
export const selectionWithoutHandler = (
// @ts-expect-error a controlled `rowSelection` needs `onRowSelectionChange` or selection changes are dropped
<DataTable data={data} columns={columns} rowSelection={rowSelection} />
);

View file

@ -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(<DataTable data={[]} columns={filterableColumns} filterMode="server" />)).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(<DataTable data={[]} columns={nameCellColumns} sortingMode="server" />)).toThrow(
/sortingMode='server'/,
);
spy.mockRestore();
});
it("throws when server pagination is missing required props", () => {
const spy = vi.spyOn(console, "error").mockImplementation(() => {});
expect(() => render(<DataTable data={[]} columns={nameCellColumns} paginationMode="server" />)).toThrow(
/paginationMode='server'/,
);
spy.mockRestore();
});
it("throws when both defaultSorting and sorting are provided", () => {
const spy = vi.spyOn(console, "error").mockImplementation(() => {});
expect(() =>
render(
<DataTable
data={[]}
columns={nameCellColumns}
defaultSorting={[{ id: "name", desc: false }]}
sorting={[{ id: "name", desc: false }]}
/>,
),
).toThrow(/defaultSorting/);
spy.mockRestore();
});
});

View file

@ -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<TData extends RowData, TValue>(
props: DataTableProps<TData, TValue>,
): 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<TData, TValue>(column: ColumnDef<TData, TValue>): string | undefined {
if ("id" in column && typeof column.id === "string") {
return column.id;
@ -442,7 +409,9 @@ function useControllable<T>(
return { value: internal, onChange: setInternal };
}
function useDataTableInstance<TData extends RowData, TValue>(props: DataTableProps<TData, TValue>): Table<TData> {
function useDataTableInstance<TData extends RowData, TValue>(
props: DataTableResolvedProps<TData, TValue>,
): Table<TData> {
const {
data,
columns,
@ -532,14 +501,7 @@ function useDataTableInstance<TData extends RowData, TValue>(props: DataTablePro
}
export function DataTable<TData extends RowData, TValue>(props: DataTableProps<TData, TValue>) {
// Validate once at construction so a misconfig surfaces immediately instead of on every render.
useState<null>(() => {
const errors = validateDataTableConfig(props);
if (errors.length > 0) {
throw new DataTableConfigError(errors);
}
return null;
});
const resolved: DataTableResolvedProps<TData, TValue> = props;
const {
isLoading = false,
@ -559,9 +521,9 @@ export function DataTable<TData extends RowData, TValue>(props: DataTableProps<T
toolbar,
paginationSlot,
footer,
} = props;
} = resolved;
const table = useDataTableInstance(props);
const table = useDataTableInstance(resolved);
const rows = table.getRowModel().rows;
const visibleColumnCount = table.getVisibleLeafColumns().length;

View file

@ -4,7 +4,7 @@ import userEvent from "@testing-library/user-event";
import { useState } from "react";
import { describe, expect, it } from "vitest";
import { createSelectionColumn, DataTable, validateDataTableConfig } from "./index";
import { createSelectionColumn, DataTable } from "./index";
interface Model {
id: string;
@ -122,16 +122,4 @@ describe("DataTable row selection", () => {
await user.click(rowBox("m1"));
expect(selectedCount()).toHaveTextContent("1");
});
it("rejects controlled rowSelection without onRowSelectionChange", () => {
const errors = validateDataTableConfig<Model, unknown>({ 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<Model, unknown>({ data, columns })).toHaveLength(0);
});
});

View file

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

View file

@ -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<TData extends RowData, TValue> {
export interface DataTableResolvedProps<TData extends RowData, TValue> {
data: TData[];
columns: ColumnDef<TData, TValue>[];
getRowId?: (row: TData, index: number, parent?: Row<TData>) => string;
@ -81,3 +81,73 @@ export interface DataTableProps<TData extends RowData, TValue> {
paginationSlot?: (table: Table<TData>) => React.ReactNode;
footer?: (table: Table<TData>) => React.ReactNode;
}
type DataTableBaseProps<TData extends RowData, TValue> = Omit<
DataTableResolvedProps<TData, TValue>,
| "sortingMode"
| "sorting"
| "onSortingChange"
| "defaultSorting"
| "paginationMode"
| "pagination"
| "onPaginationChange"
| "rowCount"
| "filterMode"
| "columnFilters"
| "onColumnFiltersChange"
| "defaultColumnFilters"
| "rowSelection"
| "onRowSelectionChange"
>;
type SortingProps =
| {
sorting: SortingState;
onSortingChange: OnChangeFn<SortingState>;
sortingMode?: SortingMode;
defaultSorting?: never;
}
| {
sortingMode?: Exclude<SortingMode, "server">;
sorting?: never;
onSortingChange?: never;
defaultSorting?: SortingState;
};
type PaginationProps =
| {
paginationMode: "server";
pagination: PaginationState;
onPaginationChange: OnChangeFn<PaginationState>;
rowCount: number;
}
| {
paginationMode?: Exclude<PaginationMode, "server">;
pagination?: PaginationState;
onPaginationChange?: OnChangeFn<PaginationState>;
rowCount?: number;
};
type FilterProps =
| {
columnFilters: ColumnFiltersState;
onColumnFiltersChange: OnChangeFn<ColumnFiltersState>;
filterMode?: FilterMode;
defaultColumnFilters?: never;
}
| {
filterMode?: Exclude<FilterMode, "server">;
columnFilters?: never;
onColumnFiltersChange?: never;
defaultColumnFilters?: ColumnFiltersState;
};
type RowSelectionProps =
| { rowSelection: RowSelectionState; onRowSelectionChange: OnChangeFn<RowSelectionState> }
| { rowSelection?: never; onRowSelectionChange?: OnChangeFn<RowSelectionState> };
export type DataTableProps<TData extends RowData, TValue> = DataTableBaseProps<TData, TValue> &
SortingProps &
PaginationProps &
FilterProps &
RowSelectionProps;

View file

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

View file

@ -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<Team[], Error> = {
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,

View file

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

View file

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

View file

@ -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<string, readonly string[]>;
export type Capability = keyof typeof CAPABILITY_ROLES;

View file

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

View file

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

View file

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