mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
commit
d46ef9aeb4
36 changed files with 1957 additions and 163 deletions
5
.github/workflows/test-litellm-ui-unit.yml
vendored
5
.github/workflows/test-litellm-ui-unit.yml
vendored
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -3399,6 +3399,8 @@ all_litellm_params = (
|
|||
+ [
|
||||
"metadata",
|
||||
"litellm_metadata",
|
||||
"keepalive_seconds",
|
||||
"allow_client_keepalive_override",
|
||||
"litellm_trace_id",
|
||||
"litellm_request_debug",
|
||||
"guardrails",
|
||||
|
|
|
|||
|
|
@ -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`)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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": [
|
||||
|
|
|
|||
|
|
@ -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 .",
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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];
|
||||
|
|
|
|||
|
|
@ -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", () => {
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
},
|
||||
],
|
||||
},
|
||||
],
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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} />
|
||||
);
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
14
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
14
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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 */
|
||||
|
|
|
|||
|
|
@ -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", () => {
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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: {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue