From fdb1ce8d782474be6541e061912be9859054a1f5 Mon Sep 17 00:00:00 2001 From: moe-berri Date: Wed, 7 Oct 2026 17:31:48 -0700 Subject: [PATCH 01/31] fix(lens): preserve numeric tags and report OTLP error codes (#45206) --- litellm-rust/crates/lens/src/ingest.rs | 77 +++++++++++++++++++++++++- litellm/tracing/types.py | 2 +- tests/unit/tracing/test_exporter.py | 51 +++++++++++++++++ 3 files changed, 127 insertions(+), 3 deletions(-) diff --git a/litellm-rust/crates/lens/src/ingest.rs b/litellm-rust/crates/lens/src/ingest.rs index a490e9ee31e..1e890f87ec2 100644 --- a/litellm-rust/crates/lens/src/ingest.rs +++ b/litellm-rust/crates/lens/src/ingest.rs @@ -46,6 +46,15 @@ pub fn response(content_type: Option<&str>, outcome: Result<(), Error>) -> Respo .map(|_| StatusCode::OK) .unwrap_or_else(|error| error.status()); let message = status.canonical_reason().unwrap_or("Trace request failed"); + let rpc_code = match status { + StatusCode::OK => 0, + StatusCode::BAD_REQUEST => 3, + StatusCode::UNAUTHORIZED => 16, + StatusCode::PAYLOAD_TOO_LARGE | StatusCode::TOO_MANY_REQUESTS => 8, + StatusCode::CONFLICT => 10, + StatusCode::SERVICE_UNAVAILABLE => 14, + _ => 2, + }; let protobuf = content_type.is_some_and(|value| { value .split(';') @@ -58,7 +67,7 @@ pub fn response(content_type: Option<&str>, outcome: Result<(), Error>) -> Respo Vec::new() } else { OtlpError { - code: 0, + code: rpc_code, message: message.into(), } .encode_to_vec() @@ -70,7 +79,7 @@ pub fn response(content_type: Option<&str>, outcome: Result<(), Error>) -> Respo if outcome.is_ok() { b"{}".to_vec() } else { - serde_json::json!({"code": 0, "message": message}) + serde_json::json!({"code": rpc_code, "message": message}) .to_string() .into_bytes() }, @@ -169,3 +178,67 @@ async fn store( .await?; Ok(()) } + +#[cfg(test)] +mod tests { + use super::{OtlpError, response}; + use crate::Error; + use axum::body::to_bytes; + use prost::Message; + use rstest::rstest; + + #[rstest] + #[case::invalid_request(Error::InvalidRequest, 400, 3, false)] + #[case::unauthenticated(Error::Unauthorized, 401, 16, false)] + #[case::payload_too_large(Error::TooLarge, 413, 8, false)] + #[case::credentials_pending(Error::CredentialsPending, 429, 8, true)] + #[case::storage_unavailable(Error::Unavailable, 503, 14, true)] + #[case::conflict(Error::TraceChanged, 409, 10, false)] + #[tokio::test] + async fn rejected_batches_have_matching_http_and_rpc_errors( + #[case] error: Error, + #[case] http_status: u16, + #[case] rpc_code: i32, + #[case] retryable: bool, + #[values("application/json", "application/x-protobuf")] content_type: &str, + ) { + let reply = response(Some(content_type), Err(error)); + assert_eq!(reply.status().as_u16(), http_status); + assert_eq!(reply.headers()["content-type"], content_type); + assert_eq!( + reply + .headers() + .get("retry-after") + .map(|v| v.to_str().unwrap()), + retryable.then_some("5") + ); + let message = reply.status().canonical_reason().unwrap(); + let body = to_bytes(reply.into_body(), 1024).await.unwrap(); + if content_type == "application/x-protobuf" { + let status = OtlpError::decode(body).unwrap(); + assert_eq!(status.code, rpc_code); + assert_eq!(status.message, message); + } else { + let status: serde_json::Value = serde_json::from_slice(&body).unwrap(); + assert_eq!( + status, + serde_json::json!({"code": rpc_code, "message": message}) + ); + } + } + + #[rstest] + #[case::json("application/json", b"{}")] + #[case::protobuf("application/x-protobuf", b"")] + #[tokio::test] + async fn accepted_batches_keep_the_empty_export_response( + #[case] content_type: &str, + #[case] expected: &[u8], + ) { + let reply = response(Some(content_type), Ok(())); + assert_eq!(reply.status(), 200); + assert_eq!(reply.headers()["content-type"], content_type); + assert!(!reply.headers().contains_key("retry-after")); + assert_eq!(to_bytes(reply.into_body(), 1024).await.unwrap(), expected); + } +} diff --git a/litellm/tracing/types.py b/litellm/tracing/types.py index 809f264d371..43490f285c2 100644 --- a/litellm/tracing/types.py +++ b/litellm/tracing/types.py @@ -49,7 +49,7 @@ class SpendLogPayload(TypedDict, total=False): cache_hit: ReadOnly[bool | None] session_id: ReadOnly[str | None] trace_id: ReadOnly[str | None] - request_tags: ReadOnly[Sequence[str] | None] + request_tags: ReadOnly[Sequence[object] | None] messages: ReadOnly[object] response: ReadOnly[object] diff --git a/tests/unit/tracing/test_exporter.py b/tests/unit/tracing/test_exporter.py index 99e9a2c9842..5f67fbbb25b 100644 --- a/tests/unit/tracing/test_exporter.py +++ b/tests/unit/tracing/test_exporter.py @@ -9,6 +9,57 @@ import pytest from litellm.tracing.exporter import MAX_BUFFER_EVENTS, MAX_EVENT_BYTES, ExportFailure, LensExporter, encode_record +@pytest.mark.asyncio +@pytest.mark.parametrize("failed", (False, True), ids=("success-event", "failure-event")) +@pytest.mark.parametrize( + ("tags", "expected"), + ( + pytest.param((7,), ("7",), id="numeric"), + pytest.param(("env:prod", 0, -7, 2.5, True, None), ("env:prod", "0", "-7", "2.5", "True", "None"), id="mixed"), + pytest.param(("env:prod", "agent:research"), ("env:prod", "agent:research"), id="strings"), + ), +) +async def test_request_tags_are_normalized_without_dropping_the_record( + failed: bool, tags: tuple[object, ...], expected: tuple[str, ...] +) -> None: + requests: Final = asyncio.Queue[httpx.Request]() + + def accept(request: httpx.Request) -> httpx.Response: + requests.put_nowait(request) + return httpx.Response(204) + + async with httpx.AsyncClient(base_url="http://lens", transport=httpx.MockTransport(accept)) as client: + exporter: Final = LensExporter(client) + exporter.start() + callback: Final = exporter.async_log_failure_event if failed else exporter.async_log_success_event + await callback( + { + "response_cost": 0.12, + "standard_logging_object": { + "id": "tagged-request", + "status": "failure" if failed else "success", + "response_cost": 0.12, + "request_tags": list(tags), + }, + }, + None, + None, + None, + ) + await exporter.aclose() + assert exporter.rows_written == 1 + assert exporter.rows_dropped == 0 + assert requests.qsize() == 1 + request: Final = requests.get_nowait() + rows: Final = json.loads(request.content) + assert request.url.path == "/internal/spend" + assert len(rows) == 1 + assert rows[0]["request_id"] == "tagged-request" + assert rows[0]["status"] == ("failure" if failed else "success") + assert rows[0]["spend"] == 0.12 + assert rows[0]["request_tags"] == list(expected) + + @pytest.mark.asyncio async def test_request_export_ignores_unrelated_model_metadata_and_preserves_billing() -> None: received: Final = asyncio.Future[httpx.Request]() From bf9f59d76b0a917858b2f618676a7a7240029fac Mon Sep 17 00:00:00 2001 From: RachelHuangZW Date: Wed, 7 Oct 2026 20:34:38 -0400 Subject: [PATCH 02/31] fix(scheduler): remove a request's queue entry once it stops waiting (#43061) * fix(scheduler): remove a request's queue entry once it stops waiting Requests admitted while a healthy deployment existed never left the priority queue, and neither did requests cancelled while waiting. The stale entries blocked later requests during cooldown and, with Redis, made add_request raise a TypeError on queues read back as JSON lists Both scheduling paths now share one polling helper that removes the entry in a finally block, whether the request was admitted, timed out or cancelled Related to #43059 * test(router): allowlist _wait_for_scheduler_turn in the router coverage check The coverage script only counts direct calls in test files. The helper is exercised through prioritized acompletion and atext_completion in test_router.py, like the other allowlisted entries * fix(scheduler): admit healthy requests before reading the queue, and clean up cancelled enqueues poll() raised on an empty queue before checking for healthy deployments. With the cleanup now rewriting the queue after every admission, a concurrent write from another replica can erase a waiting request's entry, and that request then failed while a deployment was healthy. poll() now admits as soon as a deployment is healthy and only reads the queue during cooldown add_request also moved inside the try block, so a request cancelled while its queue write is in flight still has its entry removed * refactor(scheduler): move the wait loop into Scheduler.wait_for_turn The router passes a healthy-deployments callable into the scheduler, so tests inject a Scheduler directly instead of replacing the router's scheduler attribute. The cancelled-mid-enqueue test moves to test_scheduler.py with the other scheduler tests, and poll() takes the deployments as a Sequence since it only checks whether any are healthy * test(scheduler): move scheduler tests into tests/unit * fix(scheduler): finish the queue removal when cancellation is delivered again during cleanup A second cancel, or an anyio cancel scope that re-cancels on every await, interrupted remove_request mid-write and left the entry in Redis. * fix(scheduler): read queue entries back from redis as tuples * fix(scheduler): admit a request whose queue entry vanished during cooldown * fix(scheduler): re-enqueue a request whose queue entry vanished during cooldown * test(integration): cover priority scheduler queue cleanup across instances Adds tests/integration/routing/test_priority_scheduler_queue_cleanup.py: 27 cells against the real proxy (two workers, Postgres, Redis, a scripted upstream) for every prioritized surface (/v1/chat/completions, /v1/completions, /queue/chat/completions, streaming and not, OpenAI SDK sync and async, raw httpx), the non-integer priority pass-through, the in-memory queue on one proxy, two proxies sharing a Redis queue (served requests leave no entry, a dead replica's entry is skipped or expires, a waiter that times out or disconnects removes only itself), a Redis outage mid burst and a SIGKILLed worker. Each cell asserts the caller's response, the upstream's requests by marker and the Redis queue contents. On the merge base the cross-instance cells fail with the list-of-lists TypeError and the in-memory cell with 408s behind the leaked entry; on this branch every cell passes --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/router.py | 143 +-- litellm/scheduler.py | 84 +- .../router_code_coverage.py | 1 + .../test_priority_scheduler_queue_cleanup.py | 860 ++++++++++++++++++ tests/unit/test_router/test_router.py | 59 ++ tests/unit/test_scheduler.py | 193 +++- 6 files changed, 1203 insertions(+), 137 deletions(-) create mode 100644 tests/integration/routing/test_priority_scheduler_queue_cleanup.py diff --git a/litellm/router.py b/litellm/router.py index 7b89d5eb5ac..d6730db8748 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -4534,61 +4534,21 @@ class Router: stream=False, **kwargs, ): - parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) - ### FLOW ITEM ### - _request_id: Final = str(uuid.uuid4()) - item: Final = FlowItem( - priority=priority, # 👈 SET PRIORITY FOR REQUEST - request_id=_request_id, # 👈 SET REQUEST ID - model_name=model, # 👈 SAME as 'Router' + await self._wait_for_scheduler_turn( + model=model, priority=priority, parent_otel_span=get_parent_otel_span_from_kwargs(kwargs) ) - ### [fin] ### - - ## ADDS REQUEST TO QUEUE ## - await self.scheduler.add_request(request=item) - - ## POLL QUEUE - end_time: Final = time.monotonic() + self.timeout - curr_time = time.monotonic() - poll_interval: Final = self.scheduler.polling_interval # poll every 3ms - make_request = False - - while curr_time < end_time: - _healthy_deployments, _ = await self._async_get_healthy_deployments( - model=model, parent_otel_span=parent_otel_span - ) - make_request = await self.scheduler.poll( ## POLL QUEUE ## - returns 'True' if there's healthy deployments OR if request is at top of queue - id=item.request_id, - model_name=item.model_name, - health_deployments=_healthy_deployments, - ) - if make_request: ## IF TRUE -> MAKE REQUEST - break - else: ## ELSE -> loop till default_timeout - await asyncio.sleep(poll_interval) - curr_time = time.monotonic() - - if make_request: - try: - _response: Final = await self.acompletion(model=model, messages=messages, stream=stream, **kwargs) - response_hidden_params: Final = get_hidden_params(_response) - if response_hidden_params is not None: - additional_headers: Final = cast( # cast-ok: router headers are stored as a mutable mapping - dict[str, object], response_hidden_params.setdefault("additional_headers", {}) - ) - additional_headers.update({"x-litellm-request-prioritization-used": True}) - return _response - except Exception as e: - setattr(e, "priority", priority) - raise e - else: - # Clean up the request from the scheduler queue also before raising the timeout exception - await self.scheduler.remove_request(request_id=item.request_id, model_name=item.model_name) - raise litellm.Timeout( - message="Request timed out while polling queue", - model=model, - llm_provider="openai", - ) + try: + _response: Final = await self.acompletion(model=model, messages=messages, stream=stream, **kwargs) + response_hidden_params: Final = get_hidden_params(_response) + if response_hidden_params is not None: + additional_headers: Final = cast( # cast-ok: router headers are stored as a mutable mapping + dict[str, object], response_hidden_params.setdefault("additional_headers", {}) + ) + additional_headers.update({"x-litellm-request-prioritization-used": True}) + return _response + except Exception as e: + setattr(e, "priority", priority) + raise e async def _schedule_factory( self, @@ -4598,61 +4558,32 @@ class Router: args: tuple[object, ...], kwargs: dict[str, object], ): - parent_otel_span: Final = get_parent_otel_span_from_kwargs(kwargs) - ### FLOW ITEM ### - _request_id: Final = str(uuid.uuid4()) - item: Final = FlowItem( - priority=priority, # 👈 SET PRIORITY FOR REQUEST - request_id=_request_id, # 👈 SET REQUEST ID - model_name=model, # 👈 SAME as 'Router' + await self._wait_for_scheduler_turn( + model=model, priority=priority, parent_otel_span=get_parent_otel_span_from_kwargs(kwargs) ) - ### [fin] ### + try: + _response: Final = await original_function(*args, **kwargs) + response_hidden_params: Final = get_hidden_params(_response) + if response_hidden_params is not None: + additional_headers: Final = cast( # cast-ok: router headers are stored as a mutable mapping + dict[str, object], response_hidden_params.setdefault("additional_headers", {}) + ) + additional_headers.update({"x-litellm-request-prioritization-used": True}) + return _response + except Exception as e: + setattr(e, "priority", priority) + raise e - ## ADDS REQUEST TO QUEUE ## - await self.scheduler.add_request(request=item) + async def _wait_for_scheduler_turn(self, model: str, priority: int, parent_otel_span: Span | None) -> None: + async def healthy_deployments() -> Sequence[object]: + deployments, _ = await self._async_get_healthy_deployments(model=model, parent_otel_span=parent_otel_span) + return deployments - ## POLL QUEUE - end_time: Final = time.monotonic() + self.timeout - curr_time = time.monotonic() - poll_interval: Final = self.scheduler.polling_interval # poll every 3ms - make_request = False - - while curr_time < end_time: - _healthy_deployments, _ = await self._async_get_healthy_deployments( - model=model, parent_otel_span=parent_otel_span - ) - make_request = await self.scheduler.poll( ## POLL QUEUE ## - returns 'True' if there's healthy deployments OR if request is at top of queue - id=item.request_id, - model_name=item.model_name, - health_deployments=_healthy_deployments, - ) - if make_request: ## IF TRUE -> MAKE REQUEST - break - else: ## ELSE -> loop till default_timeout - await asyncio.sleep(poll_interval) - curr_time = time.monotonic() - - if make_request: - try: - _response: Final = await original_function(*args, **kwargs) - response_hidden_params: Final = get_hidden_params(_response) - if response_hidden_params is not None: - additional_headers: Final = cast( # cast-ok: router headers are stored as a mutable mapping - dict[str, object], response_hidden_params.setdefault("additional_headers", {}) - ) - additional_headers.update({"x-litellm-request-prioritization-used": True}) - return _response - except Exception as e: - setattr(e, "priority", priority) - raise e - else: - # Clean up the request from the scheduler queue also before raising the timeout exception - await self.scheduler.remove_request(request_id=item.request_id, model_name=item.model_name) - raise litellm.Timeout( - message="Request timed out while polling queue", - model=model, - llm_provider="openai", - ) + await self.scheduler.wait_for_turn( + request=FlowItem(priority=priority, request_id=str(uuid.uuid4()), model_name=model), + timeout=self.timeout, + get_healthy_deployments=healthy_deployments, + ) def _is_prompt_management_model(self, model: str) -> bool: model_list: Final = self.get_model_list(model_name=model) diff --git a/litellm/scheduler.py b/litellm/scheduler.py index c88cc4ce0c0..73196d24e4d 100644 --- a/litellm/scheduler.py +++ b/litellm/scheduler.py @@ -1,14 +1,22 @@ +import asyncio import enum import heapq -from typing import Final +import time +from collections.abc import Awaitable, Callable, Sequence +from typing import Final, TypeAlias + +from pydantic import TypeAdapter from litellm import print_verbose from litellm._internal_context import with_service_target from litellm.caching.caching import DualCache, RedisCache from litellm.constants import DEFAULT_IN_MEMORY_TTL, DEFAULT_POLLING_INTERVAL +from litellm.exceptions import Timeout from litellm.types.llms.base import LiteLLMBaseModel SCHEDULER_QUEUE_TARGET: Final = "scheduler_queue" +QueueEntry: TypeAlias = tuple[int, str] +_QUEUE_ENTRIES: Final = TypeAdapter(list[QueueEntry]) class SchedulerCacheKeys(enum.Enum): @@ -33,7 +41,7 @@ class Scheduler: """ polling_interval: float or null - frequency of polling queue. Default is 3ms. """ - self.queue: list = [] + self.queue: list[QueueEntry] = [] default_in_memory_ttl: float | None = None if redis_cache is not None: # if redis-cache available frequently poll that instead of using in-memory. @@ -51,7 +59,7 @@ class Scheduler: # save the queue await self.save_queue(queue=queue, model_name=request.model_name) - async def poll(self, id: str, model_name: str, health_deployments: list) -> bool: + async def poll(self, request: FlowItem, health_deployments: Sequence[object]) -> bool: """ Return if request can be processed. @@ -62,30 +70,48 @@ class Scheduler: - False: * If no healthy deployments available * AND request not at the top of queue + + A request the queue no longer holds (its cache key expired or a concurrent writer erased the entry) + is put back at its priority so it keeps its place in the order instead of failing or jumping ahead """ - queue: Final = await self.get_queue(model_name=model_name) - if not queue: - raise Exception(f"Incorrectly setup. Queue is invalid. Queue={queue}") - - # ------------ - # Setup values - # ------------ - print_verbose(f"len(health_deployments): {len(health_deployments)}") - if len(health_deployments) == 0: - print_verbose(f"queue: {queue}, seeking id={id}") - # Check if the id is at the top of the heap - if queue[0][1] == id: - # Remove the item from the queue - heapq.heappop(queue) - await self.save_queue(queue=queue, model_name=model_name) - print_verbose(f"Popped id: {id}") - return True - else: - return False + if len(health_deployments) > 0: + return True + queue: Final = await self.get_queue(model_name=request.model_name) + entry: Final = (request.priority, request.request_id) + print_verbose(f"queue: {queue}, seeking {entry}") + if entry not in queue: + print_verbose(f"queue no longer holds {entry}, re-enqueueing it") + heapq.heappush(queue, entry) + if queue[0] != entry: + await self.save_queue(queue=queue, model_name=request.model_name) + return False + elif queue[0] != entry: + return False + + heapq.heappop(queue) + await self.save_queue(queue=queue, model_name=request.model_name) + print_verbose(f"Popped id: {request.request_id}") return True + async def wait_for_turn( + self, + request: FlowItem, + timeout: float, + get_healthy_deployments: Callable[[], Awaitable[Sequence[object]]], + ) -> None: + try: + await self.add_request(request=request) + end_time: Final = time.monotonic() + timeout + while time.monotonic() < end_time: + if await self.poll(request=request, health_deployments=await get_healthy_deployments()): + return + await asyncio.sleep(self.polling_interval) + finally: + await asyncio.shield(self.remove_request(request_id=request.request_id, model_name=request.model_name)) + raise Timeout(message="Request timed out while polling queue", model=request.model_name, llm_provider="openai") + async def remove_request(self, request_id: str, model_name: str) -> None: """ Remove a specific request from the priority queue for a model. @@ -118,21 +144,23 @@ class Scheduler: return self.queue @with_service_target(SCHEDULER_QUEUE_TARGET) - async def get_queue(self, model_name: str) -> list: + async def get_queue(self, model_name: str) -> list[QueueEntry]: """ - Return a queue for that specific model group + Return a queue for that specific model group. + + Redis hands the queue back as JSON lists, so every entry is validated into the + (priority, request_id) tuple the heap operations compare against. """ if self.cache is not None: _cache_key: Final = f"{SchedulerCacheKeys.queue.value}:{model_name}" response: Final = await self.cache.async_get_cache(key=_cache_key) - if response is None or not isinstance(response, list): + if not isinstance(response, list): return [] - elif isinstance(response, list): - return response + return _QUEUE_ENTRIES.validate_python(response) return self.queue @with_service_target(SCHEDULER_QUEUE_TARGET) - async def save_queue(self, queue: list, model_name: str) -> None: + async def save_queue(self, queue: list[QueueEntry], model_name: str) -> None: """ Save the updated queue of the model group """ diff --git a/tests/code_coverage_tests/router_code_coverage.py b/tests/code_coverage_tests/router_code_coverage.py index 07521fef965..cbc09edc357 100644 --- a/tests/code_coverage_tests/router_code_coverage.py +++ b/tests/code_coverage_tests/router_code_coverage.py @@ -89,6 +89,7 @@ ignored_function_names = [ "_resolve_claude_code_session_router", # Tested through Claude Code session routing in test_router.py "_get_claude_code_session_router_binding", # Tested through the two-worker session routing test in test_router.py "_apply_updated_routing_strategy_args", # Tested via update_settings in test_lowest_latency.py (file lacks "router" in name) + "_wait_for_scheduler_turn", # Tested through prioritized acompletion and atext_completion in test_router.py "arm_routing_read_prefetch", # Tested in tests/unit/caching/test_request_redis_batch_pre_call.py (file lacks "router" in name) "_configured_model_info", # Tested through get_configured_service_tiers in test_router.py "_routable_deployments", # Tested through get_configured_service_tiers and get_routable_upstream_model in test_router.py diff --git a/tests/integration/routing/test_priority_scheduler_queue_cleanup.py b/tests/integration/routing/test_priority_scheduler_queue_cleanup.py new file mode 100644 index 00000000000..a4606554d34 --- /dev/null +++ b/tests/integration/routing/test_priority_scheduler_queue_cleanup.py @@ -0,0 +1,860 @@ +from __future__ import annotations + +import asyncio +import http.client +import json +import re +import signal +import threading +import time +import uuid +from collections.abc import Callable, Iterator, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor +from contextlib import ExitStack +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from typing import Final, Literal, TypeVar +from urllib.parse import urlsplit + +import httpx +import openai +import psutil +import pytest +import yaml +from integration._support.client import Gateway, eventually, gateway_from_environment +from integration._support.database import read_rows +from integration._support.openai_wire import answering_model_discovery, chat_reply, openai_error, responses_reply +from integration._support.process import graceful_stop_seconds, owned_proxy_process +from integration._support.redis_process import OwnedRedis, owned_redis +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter +from redis import Redis + +T = TypeVar("T") +Endpoint = Literal["chat", "chat-stream", "completions", "completions-stream", "queue", "queue-stream"] +Instance = Literal["first", "second"] + +MARKER: Final = re.compile(r"sched-[0-9a-f]{32}") +UNKNOWN_MODEL: Final = "Invalid model name passed in model=" +NO_DEPLOYMENTS: Final = "No deployments available" +QUEUE_TIMEOUT: Final = "Request timed out while polling queue" +ROUTER_TIMEOUT_SECONDS: Final = 8 +PROMPT_SECONDS: Final = 4 +COOLDOWN_SECONDS: Final = 60 +USAGE: Final[dict[str, JsonValue]] = {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8} +STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +OWNED_CELL_TIMEOUT: Final = 2 * graceful_stop_seconds() + 120 +PAIR_CELL_TIMEOUT: Final = 3 * graceful_stop_seconds() + 120 +WORKER_HEALTHCHECK_ARGUMENTS: Final = ("--timeout_worker_healthcheck", str(int(graceful_stop_seconds()))) +PINNED_CONNECTION_TIMEOUT_SECONDS: Final = 30 +PINNED_LIMITS: Final = httpx.Limits(max_connections=1, max_keepalive_connections=1, keepalive_expiry=30) +ENDPOINTS: Final[tuple[Endpoint, ...]] = ( + "chat", + "chat-stream", + "completions", + "completions-stream", + "queue", + "queue-stream", +) +PAIR_GROUPS: Final = tuple(f"sched-pair-{row}" for row in ("p1", "p2", "p3", "p4", "p5", "p6", "p7", "p8")) +INMEM_GROUP: Final = "sched-inmem" +OUTAGE_GROUP: Final = "sched-outage" +KILL_GROUP: Final = "sched-kill" +GHOST_ENTRY: Final[list[JsonValue]] = [1, "ghost"] +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +JSON_LIST: Final = TypeAdapter(list[JsonValue]) + + +def new_marker() -> str: + return f"sched-{uuid.uuid4().hex}" + + +def marker_in(text: str) -> str: + found: Final = MARKER.search(text) + assert found is not None, text[:300] + return found.group(0) + + +def markers_in(text: str) -> frozenset[str]: + return frozenset(MARKER.findall(text)) + + +def upstream_name(group: str) -> str: + return f"{group}-upstream" + + +def queue_key(group: str) -> str: + return f"scheduler:queue:{group}" + + +def data_frame(frame: Mapping[str, JsonValue]) -> bytes: + return b"data: " + json.dumps(frame).encode() + b"\n\n" + + +def text_completion_reply(marker: str, model: str, *, stream: bool) -> Reply: + choice: Final[dict[str, JsonValue]] = { + "text": f"served {marker}", + "index": 0, + "logprobs": None, + "finish_reason": "stop", + } + body: Final[dict[str, JsonValue]] = { + "id": f"cmpl-{marker}", + "object": "text_completion", + "created": 1, + "model": model, + "choices": [choice], + "usage": USAGE, + } + if not stream: + return Reply(body=json.dumps(body).encode()) + first: Final = data_frame({**body, "choices": [{**choice, "finish_reason": None}]}) + last: Final = data_frame({**body, "choices": [{**choice, "text": ""}]}) + return Reply(content_type="text/event-stream", chunks=(first, last + b"data: [DONE]\n\n")) + + +@dataclass(frozen=True, slots=True) +class Upstream: + refusing: Mapping[str, threading.Event] + held: SimpleQueue[str] + release: threading.Event + hold: frozenset[str] + + def respond(self, request: Request) -> Reply: + body: Final = JSON_OBJECT.validate_json(request.body) + model: Final = str(body["model"]) + refusal: Final = self.refusing.get(model) + if refusal is not None and refusal.is_set(): + return openai_error(401) + marker: Final = marker_in(request.body.decode()) + if model in self.hold: + self.held.put(marker) + assert self.release.wait(timeout=60), "The held burst was never released" + stream: Final = body.get("stream") is True + if request.target.endswith("/responses"): + return responses_reply(f"resp_{marker}", model, f"served {marker}", stream=stream) + if request.target.endswith("/completions") and not request.target.endswith("/chat/completions"): + return text_completion_reply(marker, model, stream=stream) + return chat_reply(f"chatcmpl-{marker}", model, f"served {marker}", stream=stream) + + +def upstream_for(groups: Sequence[str], *, hold: Sequence[str] = ()) -> Upstream: + return Upstream( + {upstream_name(group): threading.Event() for group in groups}, + SimpleQueue(), + threading.Event(), + frozenset(upstream_name(group) for group in hold), + ) + + +def received_markers(wire: Wire) -> tuple[str, ...]: + return tuple(marker_in(request.body.decode()) for request in wire.drain() if request.method == "POST") + + +def assert_once(received: Sequence[str], markers: Sequence[str]) -> None: + counts: Final = {marker: received.count(marker) for marker in markers} + assert all(count == 1 for count in counts.values()), counts + + +def assert_never(received: Sequence[str], markers: Sequence[str]) -> None: + assert not set(markers) & set(received), (markers, received) + + +def deployment_entry(group: str, wire: Wire) -> dict[str, JsonValue]: + return { + "model_name": group, + "litellm_params": { + "model": f"openai/{upstream_name(group)}", + "api_base": wire.url + "/v1", + "api_key": "synthetic-openai-key", + }, + "model_info": {"id": f"{group}-deployment"}, + } + + +def owned_config( + directory: Path, + wire: Wire, + groups: Sequence[str], + settings: Mapping[str, JsonValue], + *, + cancel_on_disconnect: bool = False, +) -> Path: + base: Final = JSON_OBJECT.validate_python( + yaml.safe_load((Path(__file__).parents[1] / "proxy_config.yaml").read_text()) + ) + general_settings: Final = base.get("general_settings", {}) + assert isinstance(general_settings, dict) + general: Final[dict[str, JsonValue]] = { + **general_settings, + **({"cancel_on_disconnect": True} if cancel_on_disconnect else {}), + } + config: Final[dict[str, JsonValue]] = { + **base, + "general_settings": general, + "router_settings": dict(settings), + "model_list": [deployment_entry(group, wire) for group in groups], + } + path: Final = directory / f"scheduler-{uuid.uuid4().hex[:8]}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def cooldown_settings(cache: OwnedRedis | None) -> dict[str, JsonValue]: + redis: Final = {} if cache is None else {"redis_host": cache.host, "redis_port": cache.port} + return {"num_retries": 0, "timeout": ROUTER_TIMEOUT_SECONDS, "cooldown_time": COOLDOWN_SECONDS, **redis} + + +def redis_settings(cache: OwnedRedis) -> dict[str, JsonValue]: + return {"num_retries": 0, "redis_host": cache.host, "redis_port": cache.port} + + +def chat_body(group: str, marker: str, **extra: JsonValue) -> dict[str, JsonValue]: + return {"model": group, "messages": [{"role": "user", "content": marker}], **extra} + + +def completion_body(group: str, marker: str, **extra: JsonValue) -> dict[str, JsonValue]: + return {"model": group, "prompt": marker, **extra} + + +@dataclass(frozen=True, slots=True) +class Answer: + status: int + text: str + headers: Mapping[str, str] + seconds: float + + @property + def identity(self) -> str: + return str(JSON_OBJECT.validate_json(self.text)["id"]) + + +def timed(send: Callable[[], httpx.Response]) -> Answer: + started: Final = time.monotonic() + response: Final = send() + return Answer(response.status_code, response.text, dict(response.headers), time.monotonic() - started) + + +def settled(send: Callable[[], T], text: Callable[[T], str]) -> T: + return eventually(send, lambda observed: UNKNOWN_MODEL not in text(observed), seconds=30) + + +def sdk_settled(call: Callable[[], T]) -> T: + def attempt() -> T | openai.APIStatusError: + try: + return call() + except openai.APIStatusError as error: + if UNKNOWN_MODEL in error.message: + return error + raise + + outcome: Final = eventually(attempt, lambda observed: not isinstance(observed, openai.APIStatusError), seconds=30) + assert not isinstance(outcome, openai.APIStatusError) + return outcome + + +def post(gateway: Gateway, path: str, body: Mapping[str, JsonValue]) -> Answer: + return settled(lambda: timed(lambda: gateway.request("POST", path, body)), lambda answer: answer.text) + + +def stream_text(gateway: Gateway, path: str, body: Mapping[str, JsonValue]) -> Answer: + def send() -> Answer: + started: Final = time.monotonic() + with gateway.client.stream( + "POST", path, json=body, headers={"Authorization": f"Bearer {gateway.key}"} + ) as response: + text: Final = "".join(response.iter_text()) + return Answer(response.status_code, text, dict(response.headers), time.monotonic() - started) + + return settled(send, lambda answer: answer.text) + + +def stream_identities(text: str) -> frozenset[str]: + frames: Final = tuple( + JSON_OBJECT.validate_json(line[len("data: ") :]) + for line in text.splitlines() + if line.startswith("data: ") and line != "data: [DONE]" + ) + assert frames, text + return frozenset(str(frame["id"]) for frame in frames) + + +def call(gateway: Gateway, endpoint: Endpoint, group: str, marker: str, priority: int = 1) -> Answer: + match endpoint: + case "chat": + return post(gateway, "/v1/chat/completions", chat_body(group, marker, priority=priority)) + case "chat-stream": + return stream_text( + gateway, "/v1/chat/completions", chat_body(group, marker, priority=priority, stream=True) + ) + case "completions": + return post(gateway, "/v1/completions", completion_body(group, marker, priority=priority)) + case "completions-stream": + return stream_text( + gateway, "/v1/completions", completion_body(group, marker, priority=priority, stream=True) + ) + case "queue": + return post(gateway, "/queue/chat/completions", chat_body(group, marker, priority=priority)) + case "queue-stream": + return stream_text( + gateway, "/queue/chat/completions", chat_body(group, marker, priority=priority, stream=True) + ) + + +def served_identity(status: int, text: str, marker: str) -> str: + assert status == 200, (status, text) + identities: Final = ( + stream_identities(text) if text.startswith("data:") else frozenset({str(JSON_OBJECT.validate_json(text)["id"])}) + ) + assert identities in ({f"chatcmpl-{marker}"}, {f"cmpl-{marker}"}), (identities, marker) + assert markers_in(text) == {marker}, text + return next(iter(identities)) + + +def assert_served(answer: Answer, marker: str) -> str: + return served_identity(answer.status, answer.text, marker) + + +def spend_row_lands(identity: str) -> None: + rows: Final = eventually( + lambda: read_rows('SELECT spend FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (identity,)), + lambda values: len(values) == 1, + seconds=70, + ) + assert rows[0]["spend"] is not None, rows + + +def sdk(gateway: Gateway) -> openai.OpenAI: + return openai.OpenAI(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0, timeout=30) + + +def queue_entries(cache: OwnedRedis, group: str) -> list[JsonValue] | None: + with Redis(host=cache.host, port=cache.port) as client: + raw: Final = client.get(queue_key(group)) + if raw is None: + return None + assert isinstance(raw, bytes), raw + return JSON_LIST.validate_json(raw) + + +def write_queue(cache: OwnedRedis, group: str, entries: Sequence[JsonValue]) -> None: + with Redis(host=cache.host, port=cache.port) as client: + client.set(queue_key(group), json.dumps(list(entries))) + + +def persist_queue(cache: OwnedRedis, group: str) -> None: + with Redis(host=cache.host, port=cache.port) as client: + client.persist(queue_key(group)) + + +def entry_priority(entry: JsonValue) -> int: + assert isinstance(entry, list), entry + priority: Final = entry[0] + assert isinstance(priority, int), entry + return priority + + +def waiting_entries(cache: OwnedRedis, group: str, priority: int) -> list[JsonValue]: + def read() -> list[JsonValue]: + return queue_entries(cache, group) or [] + + return eventually( + read, lambda found: any(entry_priority(entry) == priority for entry in found), seconds=PROMPT_SECONDS + ) + + +def cooled(gateway: Gateway, upstream: Upstream, group: str) -> None: + upstream.refusing[upstream_name(group)].set() + trip: Final = post(gateway, "/v1/chat/completions", chat_body(group, new_marker())) + assert trip.status == 401, (trip.status, trip.text) + eventually( + lambda: post(gateway, "/v1/chat/completions", chat_body(group, new_marker())), + lambda answer: answer.status == 429 and NO_DEPLOYMENTS in answer.text, + seconds=10, + ) + + +def assert_refused_at_once(answer: Answer) -> None: + assert answer.status == 429, (answer.status, answer.text) + assert NO_DEPLOYMENTS in answer.text, answer.text + assert answer.seconds < PROMPT_SECONDS, answer.seconds + + +def assert_timed_out_polling(answer: Answer) -> None: + assert answer.status == 408, (answer.status, answer.text) + assert QUEUE_TIMEOUT in answer.text, answer.text + assert answer.seconds >= ROUTER_TIMEOUT_SECONDS, answer.seconds + + +@pytest.fixture(scope="module") +def rig_upstream() -> Iterator[tuple[Upstream, Wire]]: + upstream: Final = upstream_for(()) + with wire_server(answering_model_discovery(upstream.respond)) as wire: + yield upstream, wire + + +@pytest.fixture +def rig_model(gateway: Gateway, rig_upstream: tuple[Upstream, Wire]) -> Iterator[tuple[str, Wire]]: + _, wire = rig_upstream + with gateway.scenario() as scenario: + yield scenario.model(api_base=wire.url + "/v1"), wire + + +def test_sdk_chat_with_priority_is_served_once_and_billed(gateway: Gateway, rig_model: tuple[str, Wire]) -> None: + model, wire = rig_model + marker: Final = new_marker() + client: Final = sdk(gateway) + raw: Final = sdk_settled( + lambda: client.chat.completions.with_raw_response.create( + model=model, messages=[{"role": "user", "content": marker}], extra_body={"priority": 1} + ) + ) + assert served_identity(raw.status_code, raw.text, marker) == f"chatcmpl-{marker}" + (upstream_request,) = tuple(request for request in wire.drain() if marker.encode() in request.body) + assert "priority" not in JSON_OBJECT.validate_json(upstream_request.body), upstream_request.body + spend_row_lands(f"chatcmpl-{marker}") + + +def test_sdk_chat_stream_with_priority_is_served_once_and_billed(gateway: Gateway, rig_model: tuple[str, Wire]) -> None: + model, wire = rig_model + marker: Final = new_marker() + client: Final = sdk(gateway) + chunks: Final = sdk_settled( + lambda: tuple( + client.chat.completions.create( + model=model, messages=[{"role": "user", "content": marker}], stream=True, extra_body={"priority": 1} + ) + ) + ) + assert {chunk.id for chunk in chunks} == {f"chatcmpl-{marker}"}, chunks + content: Final = "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) + assert content == f"served {marker}", chunks + assert_once(received_markers(wire), (marker,)) + spend_row_lands(f"chatcmpl-{marker}") + + +def test_async_sdk_chat_with_priority_is_served_once(gateway: Gateway, rig_model: tuple[str, Wire]) -> None: + model, wire = rig_model + marker: Final = new_marker() + + async def send() -> tuple[int, str]: + async with openai.AsyncOpenAI( + base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0, timeout=30 + ) as client: + raw: Final = await client.chat.completions.with_raw_response.create( + model=model, messages=[{"role": "user", "content": marker}], extra_body={"priority": 1} + ) + return raw.status_code, raw.text + + status, text = sdk_settled(lambda: asyncio.run(send())) + assert served_identity(status, text, marker) == f"chatcmpl-{marker}" + assert_once(received_markers(wire), (marker,)) + + +def test_raw_chat_with_duplicate_priority_keys_takes_the_last_one( + gateway: Gateway, rig_model: tuple[str, Wire] +) -> None: + model, wire = rig_model + marker: Final = new_marker() + body: Final = ( + '{"model": "%s", "messages": [{"role": "user", "content": "%s"}], "priority": "not-an-int", "priority": 1}' + % (model, marker) + ) + + def send() -> httpx.Response: + return gateway.client.post( + "/v1/chat/completions", + content=body.encode(), + headers={"Authorization": f"Bearer {gateway.key}", "Content-Type": "application/json"}, + ) + + answer: Final = settled(lambda: timed(send), lambda observed: observed.text) + assert_served(answer, marker) + (upstream_request,) = tuple(request for request in wire.drain() if marker.encode() in request.body) + assert "priority" not in JSON_OBJECT.validate_json(upstream_request.body), upstream_request.body + + +@pytest.mark.parametrize("endpoint", ("completions", "completions-stream", "queue", "queue-stream")) +def test_other_prioritized_endpoints_are_served_once_and_billed( + gateway: Gateway, rig_model: tuple[str, Wire], endpoint: Endpoint +) -> None: + model, wire = rig_model + marker: Final = new_marker() + answer: Final = call(gateway, endpoint, model, marker) + identity: Final = assert_served(answer, marker) + if endpoint == "queue": + assert answer.headers.get("x-litellm-priority") == "1", answer.headers + assert_once(received_markers(wire), (marker,)) + spend_row_lands(identity) + + +@pytest.mark.parametrize("priority", ("1", [1], "", "p" * 5120, 0), ids=("string", "list", "empty", "5kb", "zero")) +def test_non_integer_priority_bypasses_the_scheduler_and_is_forwarded( + gateway: Gateway, rig_model: tuple[str, Wire], priority: JsonValue +) -> None: + model, wire = rig_model + marker: Final = new_marker() + answer: Final = post(gateway, "/v1/chat/completions", chat_body(model, marker, priority=priority)) + assert_served(answer, marker) + (upstream_request,) = tuple(request for request in wire.drain() if marker.encode() in request.body) + assert JSON_OBJECT.validate_json(upstream_request.body).get("priority") == priority + + +def test_unauthenticated_prioritized_request_never_reaches_the_upstream( + gateway: Gateway, rig_model: tuple[str, Wire] +) -> None: + model, wire = rig_model + marker: Final = new_marker() + response: Final = gateway.client.post("/v1/chat/completions", json=chat_body(model, marker, priority=1)) + assert response.status_code == 401, response.text + assert_never(received_markers(wire), (marker,)) + + +@pytest.mark.parametrize("path", ("/v1/messages", "/v1/responses"), ids=("messages", "responses")) +def test_priority_on_unscheduled_routes_is_served_once( + gateway: Gateway, rig_model: tuple[str, Wire], path: str +) -> None: + model, wire = rig_model + marker: Final = new_marker() + body: Final[dict[str, JsonValue]] = ( + {"model": model, "max_tokens": 32, "messages": [{"role": "user", "content": marker}], "priority": 1} + if path == "/v1/messages" + else {"model": model, "input": marker, "priority": 1} + ) + answer: Final = post(gateway, path, body) + assert answer.status == 200, (answer.status, answer.text) + assert markers_in(answer.text) == {marker}, answer.text + assert_once(received_markers(wire), (marker,)) + + +def pinned(gateway: Gateway, stack: ExitStack) -> Gateway: + client: Final = stack.enter_context( + httpx.Client(base_url=gateway.client.base_url, timeout=15, trust_env=False, limits=PINNED_LIMITS) + ) + return Gateway(client, gateway.key, gateway.upstream_url) + + +@pytest.mark.timeout(OWNED_CELL_TIMEOUT) +def test_in_memory_queue_forgets_served_requests_before_a_cooldown(gateway: Gateway, tmp_path: Path) -> None: + upstream: Final = upstream_for((INMEM_GROUP,)) + with ExitStack() as stack: + wire: Final = stack.enter_context(wire_server(answering_model_discovery(upstream.respond))) + config: Final = owned_config(tmp_path, wire, (INMEM_GROUP,), cooldown_settings(None)) + owned: Final = stack.enter_context( + owned_proxy_process( + gateway, tmp_path, {}, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS + ) + ) + worker: Final = pinned(owned.gateway, stack) + served: Final = new_marker() + assert_served(post(worker, "/v1/chat/completions", chat_body(INMEM_GROUP, served, priority=1)), served) + cooled(worker, upstream, INMEM_GROUP) + waiting: Final = (new_marker(), new_marker()) + for marker in waiting: + assert_refused_at_once(post(worker, "/v1/chat/completions", chat_body(INMEM_GROUP, marker, priority=2))) + received: Final = received_markers(wire) + assert_once(received, (served,)) + assert_never(received, waiting) + + +@dataclass(frozen=True, slots=True) +class Pair: + first: Gateway + second: Gateway + cache: OwnedRedis + wire: Wire + upstream: Upstream + + def at(self, instance: Instance) -> Gateway: + return self.first if instance == "first" else self.second + + +@pytest.fixture(scope="module") +def pair(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Pair]: + directory: Final = tmp_path_factory.mktemp("scheduler-pair") + upstream: Final = upstream_for(PAIR_GROUPS) + with ExitStack() as stack: + gateway: Final = stack.enter_context(gateway_from_environment()) + cache: Final = stack.enter_context(owned_redis(directory)) + wire: Final = stack.enter_context(wire_server(answering_model_discovery(upstream.respond))) + config: Final = owned_config(directory, wire, PAIR_GROUPS, cooldown_settings(cache), cancel_on_disconnect=True) + overrides: Final = {"REDIS_HOST": cache.host, "REDIS_PORT": str(cache.port)} + first: Final = stack.enter_context( + owned_proxy_process( + gateway, directory, overrides, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS + ) + ) + second: Final = stack.enter_context( + owned_proxy_process( + gateway, directory, overrides, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS + ) + ) + yield Pair(first.gateway, second.gateway, cache, wire, upstream) + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_served_request_leaves_no_queue_entry_for_the_other_instance(pair: Pair) -> None: + group: Final = PAIR_GROUPS[0] + markers: Final = (new_marker(), new_marker(), new_marker()) + assert_served(post(pair.first, "/v1/chat/completions", chat_body(group, markers[0], priority=1)), markers[0]) + assert queue_entries(pair.cache, group) == [] + assert_served(post(pair.second, "/v1/chat/completions", chat_body(group, markers[1], priority=1)), markers[1]) + assert_served(post(pair.first, "/v1/chat/completions", chat_body(group, markers[2], priority=1)), markers[2]) + assert_once(received_markers(pair.wire), markers) + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_concurrent_prioritized_burst_across_instances_and_endpoints_is_served(pair: Pair) -> None: + group: Final = PAIR_GROUPS[1] + primer: Final = new_marker() + assert_served(post(pair.first, "/v1/chat/completions", chat_body(group, primer, priority=1)), primer) + instances: Final[tuple[Instance, ...]] = ("first", "second") + plan: Final = tuple((instance, endpoint, new_marker()) for instance in instances for endpoint in ENDPOINTS) + + def one(item: tuple[Instance, Endpoint, str]) -> Answer: + instance, endpoint, marker = item + return call(pair.at(instance), endpoint, group, marker) + + with ThreadPoolExecutor(max_workers=len(plan)) as pool: + answers: Final = tuple(pool.map(one, plan)) + for (_, _, marker), answer in zip(plan, answers, strict=True): + assert_served(answer, marker) + closing: Final = new_marker() + assert_served(post(pair.second, "/v1/chat/completions", chat_body(group, closing, priority=1)), closing) + assert_once(received_markers(pair.wire), (primer, *(marker for _, _, marker in plan), closing)) + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_fresh_queue_during_a_cooldown_is_refused_at_once(pair: Pair) -> None: + group: Final = PAIR_GROUPS[2] + cooled(pair.second, pair.upstream, group) + client: Final = sdk(pair.second) + started: Final = time.monotonic() + with pytest.raises(openai.APIStatusError) as refused: + client.completions.create(model=group, prompt=new_marker(), extra_body={"priority": 1}) + assert refused.value.status_code == 429, refused.value.message + assert NO_DEPLOYMENTS in refused.value.message + assert time.monotonic() - started < PROMPT_SECONDS + assert queue_entries(pair.cache, group) == [] + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_served_request_on_one_instance_does_not_block_a_waiter_on_the_other(pair: Pair) -> None: + group: Final = PAIR_GROUPS[3] + served: Final = new_marker() + assert_served(post(pair.first, "/v1/chat/completions", chat_body(group, served, priority=1)), served) + cooled(pair.second, pair.upstream, group) + assert_refused_at_once(post(pair.second, "/v1/chat/completions", chat_body(group, new_marker(), priority=2))) + + +def assert_refused_before_the_router_timeout(answer: Answer) -> None: + assert answer.status == 429, (answer.status, answer.text) + assert NO_DEPLOYMENTS in answer.text, answer.text + assert answer.seconds < ROUTER_TIMEOUT_SECONDS, answer.seconds + + +def assert_drained(cache: OwnedRedis, group: str) -> None: + assert queue_entries(cache, group) in ([], None) + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_waiter_behind_an_expired_dead_replica_entry_is_re_enqueued_and_proceeds(pair: Pair) -> None: + group: Final = PAIR_GROUPS[4] + cooled(pair.second, pair.upstream, group) + write_queue(pair.cache, group, [GHOST_ENTRY]) + answer: Final = post(pair.second, "/queue/chat/completions", chat_body(group, new_marker(), priority=2)) + assert_refused_before_the_router_timeout(answer) + assert_drained(pair.cache, group) + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_waiter_behind_a_live_entry_times_out_and_removes_only_itself(pair: Pair) -> None: + group: Final = PAIR_GROUPS[5] + cooled(pair.second, pair.upstream, group) + write_queue(pair.cache, group, [GHOST_ENTRY]) + with ThreadPoolExecutor(max_workers=1) as pool: + waiter: Final = pool.submit( + post, pair.second, "/v1/chat/completions", chat_body(group, new_marker(), priority=2) + ) + waiting_entries(pair.cache, group, 2) + persist_queue(pair.cache, group) + assert_timed_out_polling(waiter.result(timeout=30)) + assert queue_entries(pair.cache, group) == [GHOST_ENTRY] + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_waiter_proceeds_once_the_entry_ahead_is_cleared(pair: Pair) -> None: + group: Final = PAIR_GROUPS[6] + cooled(pair.second, pair.upstream, group) + write_queue(pair.cache, group, [GHOST_ENTRY]) + with ThreadPoolExecutor(max_workers=1) as pool: + waiter: Final = pool.submit( + post, pair.second, "/v1/chat/completions", chat_body(group, new_marker(), priority=2) + ) + entries: Final = waiting_entries(pair.cache, group, 2) + write_queue(pair.cache, group, [entry for entry in entries if entry_priority(entry) == 2]) + assert_refused_before_the_router_timeout(waiter.result(timeout=30)) + assert_drained(pair.cache, group) + + +def pinned_connection(gateway: Gateway) -> http.client.HTTPConnection: + address: Final = urlsplit(str(gateway.client.base_url)) + assert address.hostname is not None and address.port is not None, address + return http.client.HTTPConnection(address.hostname, address.port, timeout=PINNED_CONNECTION_TIMEOUT_SECONDS) + + +def send_pinned(connection: http.client.HTTPConnection, key: str, body: Mapping[str, JsonValue]) -> None: + connection.request( + "POST", + "/v1/chat/completions", + body=json.dumps(body).encode(), + headers={"Authorization": f"Bearer {key}", "Content-Type": "application/json"}, + ) + + +def post_pinned(connection: http.client.HTTPConnection, key: str, body: Mapping[str, JsonValue]) -> tuple[int, str]: + send_pinned(connection, key, body) + response: Final = connection.getresponse() + text: Final = response.read().decode() + assert connection.sock is not None, "the proxy closed the pinned connection" + return response.status, text + + +def cooled_over(connection: http.client.HTTPConnection, key: str, upstream: Upstream, group: str) -> None: + upstream.refusing[upstream_name(group)].set() + trip: Final = post_pinned(connection, key, chat_body(group, new_marker())) + assert trip[0] == 401, trip + eventually( + lambda: post_pinned(connection, key, chat_body(group, new_marker())), + lambda answer: answer[0] == 429 and NO_DEPLOYMENTS in answer[1], + seconds=10, + ) + + +@pytest.mark.timeout(PAIR_CELL_TIMEOUT) +def test_disconnected_waiter_is_removed_from_the_queue(pair: Pair) -> None: + group: Final = PAIR_GROUPS[7] + waiter: Final = pinned_connection(pair.second) + cooled_over(waiter, pair.second.key, pair.upstream, group) + write_queue(pair.cache, group, [GHOST_ENTRY]) + send_pinned(waiter, pair.second.key, chat_body(group, new_marker(), priority=2)) + waiting_entries(pair.cache, group, 2) + persist_queue(pair.cache, group) + waiter.close() + eventually( + lambda: queue_entries(pair.cache, group), lambda entries: entries == [GHOST_ENTRY], seconds=PROMPT_SECONDS + ) + + +Outcome = Answer | httpx.TransportError + + +def burst(gateway: Gateway, group: str, count: int) -> tuple[tuple[Endpoint, str, Outcome], ...]: + plan: Final[tuple[tuple[Endpoint, str], ...]] = tuple( + (ENDPOINTS[index % len(ENDPOINTS)], new_marker()) for index in range(count) + ) + + def one(item: tuple[Endpoint, str]) -> Outcome: + endpoint, marker = item + with httpx.Client(base_url=gateway.client.base_url, timeout=60, trust_env=False) as client: + try: + return call(Gateway(client, gateway.key, gateway.upstream_url), endpoint, group, marker) + except httpx.TransportError as error: + return error + + with ThreadPoolExecutor(max_workers=count) as pool: + outcomes: Final = tuple(pool.map(one, plan)) + return tuple((endpoint, marker, outcome) for (endpoint, marker), outcome in zip(plan, outcomes, strict=True)) + + +def assert_all_served(outcomes: Sequence[tuple[Endpoint, str, Outcome]]) -> tuple[str, ...]: + for _, marker, outcome in outcomes: + assert isinstance(outcome, Answer), repr(outcome) + assert_served(outcome, marker) + return tuple(marker for _, marker, _ in outcomes) + + +def served_eventually(gateway: Gateway, group: str, marker: str) -> None: + def send() -> Outcome: + try: + return post(gateway, "/v1/chat/completions", chat_body(group, marker, priority=1)) + except httpx.TransportError as error: + return error + + answer: Final = eventually(send, lambda outcome: isinstance(outcome, Answer), seconds=30) + assert isinstance(answer, Answer) + assert_served(answer, marker) + + +@pytest.mark.timeout(OWNED_CELL_TIMEOUT) +def test_prioritized_requests_survive_a_redis_outage(gateway: Gateway, tmp_path: Path) -> None: + upstream: Final = upstream_for((OUTAGE_GROUP,)) + with owned_redis(tmp_path) as cache, wire_server(answering_model_discovery(upstream.respond)) as wire: + config: Final = owned_config(tmp_path, wire, (OUTAGE_GROUP,), redis_settings(cache)) + overrides: Final = { + "REDIS_HOST": cache.host, + "REDIS_PORT": str(cache.port), + "REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT": "1", + } + with owned_proxy_process( + gateway, tmp_path, overrides, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS + ) as owned: + before: Final = assert_all_served(burst(owned.gateway, OUTAGE_GROUP, 12)) + cache.stop() + during: Final = assert_all_served(burst(owned.gateway, OUTAGE_GROUP, 12)) + cache.start() + after: Final = assert_all_served(burst(owned.gateway, OUTAGE_GROUP, 12)) + assert_once(received_markers(wire), (*before, *during, *after)) + eventually(lambda: queue_entries(cache, OUTAGE_GROUP), lambda entries: entries in (None, []), seconds=10) + + +def established_upstream_connections(pid: int, wire: Wire) -> int: + port: Final = urlsplit(wire.url).port + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) + + +@pytest.mark.timeout(OWNED_CELL_TIMEOUT) +def test_sibling_worker_keeps_serving_prioritized_requests_after_a_worker_is_killed( + gateway: Gateway, tmp_path: Path +) -> None: + upstream: Final = upstream_for((KILL_GROUP,), hold=(KILL_GROUP,)) + with owned_redis(tmp_path) as cache, wire_server(answering_model_discovery(upstream.respond)) as wire: + config: Final = owned_config(tmp_path, wire, (KILL_GROUP,), redis_settings(cache)) + overrides: Final = {"REDIS_HOST": cache.host, "REDIS_PORT": str(cache.port)} + with owned_proxy_process( + gateway, tmp_path, overrides, config=config, workers=2, extra_arguments=WORKER_HEALTHCHECK_ARGUMENTS + ) as owned: + workers: Final = eventually( + lambda: tuple(int(found.group(1)) for found in STARTED_WORKER.finditer(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=graceful_stop_seconds(), + ) + with ThreadPoolExecutor(max_workers=1) as pool: + pending: Final = pool.submit(burst, owned.gateway, KILL_GROUP, 20) + eventually(upstream.held.qsize, lambda size: size == 20, seconds=60) + held_by: Final = {pid: established_upstream_connections(pid, wire) for pid in workers} + assert sum(held_by.values()) == 20, held_by + victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__) + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + upstream.release.set() + outcomes: Final = pending.result(timeout=90) + answered: Final = tuple(outcome for outcome in outcomes if isinstance(outcome[2], Answer)) + failed: Final = tuple(outcome for outcome in outcomes if not isinstance(outcome[2], Answer)) + assert len(answered) == held_by[survivor_pid], (held_by, len(answered)) + assert len(failed) == held_by[victim_pid], (held_by, len(failed)) + assert_all_served(answered) + follow_up: Final = new_marker() + served_eventually(owned.gateway, KILL_GROUP, follow_up) + eventually( + lambda: len(STARTED_WORKER.findall(owned.log.read_text())), + lambda count: count == 3, + seconds=graceful_stop_seconds(), + ) + assert_once(received_markers(wire), (*(marker for _, marker, _ in outcomes), follow_up)) diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 8f3133dcaa4..ec9e4861fdd 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -65,6 +65,7 @@ from litellm.router_utils.cooldown_handlers import ( ) from litellm.router_utils.fallback_event_handlers import DISABLE_FALLBACKS_METADATA_KEY from litellm.router_utils.router_callbacks.track_deployment_metrics import get_deployment_successes_for_current_minute +from litellm.scheduler import FlowItem from litellm.types.llms.openai import ChatCompletionRequest from litellm.types.router import ( CustomRoutingStrategyBase, @@ -19949,6 +19950,64 @@ def test_access_windows_filter_reserved_deployments_method(): ] == ["reserved-deployment", "open-deployment"] +def _scheduled_router(timeout: float) -> Router: + return Router( + model_list=[ + { + "model_name": "sched-model", + "litellm_params": {"model": "openai/sched-model", "api_key": "sk-fake", "mock_response": "hi"}, + "model_info": {"id": "sched-deployment"}, + } + ], + timeout=timeout, + ) + + +async def _send_scheduled_chat(router: Router, priority: int) -> object: + return await router.acompletion( + model="sched-model", messages=[{"role": "user", "content": "hi"}], priority=priority + ) + + +async def _send_scheduled_text(router: Router, priority: int) -> object: + return await router.atext_completion(model="sched-model", prompt="hi", priority=priority) + + +@pytest.mark.parametrize( + "send", [_send_scheduled_chat, _send_scheduled_text], ids=["schedule_acompletion", "schedule_factory"] +) +@pytest.mark.asyncio +async def test_admitted_prioritized_request_does_not_block_later_request_during_cooldown( + send: Callable[[Router, int], Awaitable[object]], +): + from litellm.types.router import RouterRateLimitError + + router: Final = _scheduled_router(timeout=1) + await send(router, 1) + _cool_down(router, "sched-deployment") + + with pytest.raises(RouterRateLimitError, match="cooldown"): + await send(router, 2) + + +@pytest.mark.parametrize("stop_waiting", ["cancel", "timeout"]) +@pytest.mark.asyncio +async def test_prioritized_request_leaves_queue_when_it_stops_waiting(stop_waiting: Literal["cancel", "timeout"]): + router: Final = _scheduled_router(timeout=0.5) + _cool_down(router, "sched-deployment") + await router.scheduler.add_request(FlowItem(priority=0, request_id="head-of-queue", model_name="sched-model")) + waiting: Final = asyncio.create_task(_send_scheduled_chat(router, 5)) + await asyncio.sleep(0.05) + assert len(await router.scheduler.get_queue("sched-model")) == 2 + + if stop_waiting == "cancel": + waiting.cancel() + with pytest.raises(asyncio.CancelledError if stop_waiting == "cancel" else litellm.Timeout): + await waiting + + assert await router.scheduler.get_queue("sched-model") == [(0, "head-of-queue")] + + @pytest.mark.asyncio async def test_bare_model_group_served_by_wildcard_deployment_uses_provider_prefixed_fallback_key() -> None: """Claude Code sends the bare "claude-sonnet-4-6" to /v1/messages; routing serves it through the diff --git a/tests/unit/test_scheduler.py b/tests/unit/test_scheduler.py index 553fb44cb27..e0968e9b2cf 100644 --- a/tests/unit/test_scheduler.py +++ b/tests/unit/test_scheduler.py @@ -3,7 +3,10 @@ import asyncio import importlib +import json import os +from collections.abc import Sequence +from typing import Final import pytest @@ -26,8 +29,8 @@ async def test_scheduler_diff_model_names(): await scheduler.add_request(item1) await scheduler.add_request(item2) - assert await scheduler.poll(id="10", model_name="gpt-3.5-turbo", health_deployments=[{"key": "value"}]) == True - assert await scheduler.poll(id="11", model_name="gpt-4", health_deployments=[{"key": "value"}]) == True + assert await scheduler.poll(request=item1, health_deployments=[{"key": "value"}]) == True + assert await scheduler.poll(request=item2, health_deployments=[{"key": "value"}]) == True @pytest.mark.asyncio @@ -50,7 +53,7 @@ async def test_scheduler_poll_persists_queue_to_cache(): await scheduler.add_request(item1) await scheduler.add_request(item2) - await scheduler.poll(id="10", model_name="gpt-3.5-turbo", health_deployments=[]) + await scheduler.poll(request=item1, health_deployments=[]) queue_key = f"{SchedulerCacheKeys.queue.value}:{item1.model_name}" updated_queue = redis_cache.store[queue_key] @@ -145,6 +148,190 @@ async def test_scheduler_queue_cleanup_on_timeout(): assert queue_after[0][1] == "req-0", "Expected req-0 (priority 0) to be at front" +@pytest.mark.asyncio +async def test_poll_admits_request_missing_from_queue_while_a_deployment_is_healthy(): + scheduler: Final = Scheduler() + + assert await scheduler.poll( + request=FlowItem(priority=1, request_id="erased-by-concurrent-write", model_name="sched-model"), + health_deployments=[{"model_info": {"id": "a"}}], + ) + + +@pytest.mark.asyncio +async def test_poll_during_cooldown_admits_only_the_head_of_the_queue(): + scheduler: Final = Scheduler() + later: Final = FlowItem(priority=2, request_id="later", model_name="sched-model") + head: Final = FlowItem(priority=1, request_id="head", model_name="sched-model") + await scheduler.add_request(later) + await scheduler.add_request(head) + + assert not await scheduler.poll(request=later, health_deployments=[]) + assert await scheduler.poll(request=head, health_deployments=[]) + assert await scheduler.get_queue("sched-model") == [(2, "later")] + + +@pytest.mark.asyncio +async def test_poll_during_cooldown_re_enqueues_a_request_a_concurrent_writer_erased_behind_the_head(): + scheduler: Final = Scheduler() + still_queued: Final = FlowItem(priority=0, request_id="still-queued", model_name="sched-model") + erased: Final = FlowItem(priority=1, request_id="erased-by-concurrent-write", model_name="sched-model") + await scheduler.add_request(still_queued) + + assert not await scheduler.poll(request=erased, health_deployments=[]) + assert await scheduler.get_queue("sched-model") == [(0, "still-queued"), (1, "erased-by-concurrent-write")] + assert await scheduler.poll(request=still_queued, health_deployments=[]) + assert await scheduler.poll(request=erased, health_deployments=[]) + assert await scheduler.get_queue("sched-model") == [] + + +@pytest.mark.asyncio +async def test_poll_during_cooldown_admits_an_erased_request_that_outranks_the_queue(): + scheduler: Final = Scheduler() + await scheduler.add_request(FlowItem(priority=2, request_id="still-queued", model_name="sched-model")) + + assert await scheduler.poll( + request=FlowItem(priority=0, request_id="erased-urgent", model_name="sched-model"), health_deployments=[] + ) + assert await scheduler.get_queue("sched-model") == [(2, "still-queued")] + + +class _ExpiringCache: + def __init__(self) -> None: + self.store: dict[str, object] = {} + + async def async_get_cache(self, key: str, **kwargs: object) -> object: + return self.store.get(key) + + async def async_set_cache(self, key: str, value: object, **kwargs: object) -> None: + self.store[key] = value + + +@pytest.mark.asyncio +async def test_wait_for_turn_during_cooldown_survives_the_queue_key_expiring(): + cache: Final = _ExpiringCache() + scheduler: Final = Scheduler(redis_cache=cache) + scheduler.cache.in_memory_cache.cache_dict.clear() + + async def no_healthy_deployments_after_the_key_expired() -> Sequence[object]: + cache.store.clear() + scheduler.cache.in_memory_cache.cache_dict.clear() + return () + + await scheduler.wait_for_turn( + request=FlowItem(priority=1, request_id="sole-waiter", model_name="sched-model"), + timeout=5, + get_healthy_deployments=no_healthy_deployments_after_the_key_expired, + ) + + assert await scheduler.get_queue("sched-model") == [] + + + +class _JsonRoundTripRedisCache: + def __init__(self) -> None: + self.store: dict[str, str] = {} + + async def async_get_cache(self, key: str, **kwargs: object) -> object: + raw: Final = self.store.get(key) + return None if raw is None else json.loads(raw) + + async def async_set_cache(self, key: str, value: object, **kwargs: object) -> None: + self.store[key] = json.dumps(value) + + +@pytest.mark.asyncio +async def test_second_replica_enqueues_behind_a_queue_decoded_from_redis(): + redis_cache: Final = _JsonRoundTripRedisCache() + replica_a: Final = Scheduler(redis_cache=redis_cache) + replica_b: Final = Scheduler(redis_cache=redis_cache) + waiting_on_a: Final = FlowItem(priority=1, request_id="waiting-on-a", model_name="sched-model") + urgent_on_b: Final = FlowItem(priority=0, request_id="urgent-on-b", model_name="sched-model") + await replica_a.add_request(waiting_on_a) + + await replica_b.add_request(urgent_on_b) + + assert await replica_b.get_queue("sched-model") == [(0, "urgent-on-b"), (1, "waiting-on-a")] + assert await replica_b.poll(request=urgent_on_b, health_deployments=[]) + assert await replica_b.poll(request=waiting_on_a, health_deployments=[]) + await replica_b.remove_request(request_id="waiting-on-a", model_name="sched-model") + assert await replica_b.get_queue("sched-model") == [] + +class _PausesAfterEnqueueScheduler(Scheduler): + def __init__(self) -> None: + super().__init__() + self.enqueued: Final = asyncio.Event() + + async def add_request(self, request: FlowItem) -> None: + await super().add_request(request) + self.enqueued.set() + await asyncio.Event().wait() + + +async def _no_healthy_deployments() -> Sequence[object]: + return () + + +@pytest.mark.asyncio +async def test_wait_for_turn_removes_entry_when_cancelled_mid_enqueue(): + scheduler: Final = _PausesAfterEnqueueScheduler() + waiting: Final = asyncio.create_task( + scheduler.wait_for_turn( + request=FlowItem(priority=1, request_id="cancelled", model_name="sched-model"), + timeout=5, + get_healthy_deployments=_no_healthy_deployments, + ) + ) + await scheduler.enqueued.wait() + assert await scheduler.get_queue("sched-model") == [(1, "cancelled")] + + waiting.cancel() + with pytest.raises(asyncio.CancelledError): + await waiting + + assert await scheduler.get_queue("sched-model") == [] + + +class _HeldRemovalScheduler(Scheduler): + def __init__(self) -> None: + super().__init__() + self.removing: Final = asyncio.Event() + self.finish_removal: Final = asyncio.Event() + + async def remove_request(self, request_id: str, model_name: str) -> None: + self.removing.set() + await self.finish_removal.wait() + await super().remove_request(request_id=request_id, model_name=model_name) + + +@pytest.mark.asyncio +async def test_wait_for_turn_finishes_removal_when_cancelled_again_during_cleanup(): + scheduler: Final = _HeldRemovalScheduler() + await scheduler.add_request(FlowItem(priority=0, request_id="head", model_name="sched-model")) + polling: Final = asyncio.Event() + + async def no_healthy_deployments() -> Sequence[object]: + polling.set() + return () + + waiting: Final = asyncio.create_task( + scheduler.wait_for_turn( + request=FlowItem(priority=1, request_id="cancelled", model_name="sched-model"), + timeout=5, + get_healthy_deployments=no_healthy_deployments, + ) + ) + await polling.wait() + waiting.cancel() + await scheduler.removing.wait() + waiting.cancel() + scheduler.finish_removal.set() + with pytest.raises(asyncio.CancelledError): + await waiting + + assert await scheduler.get_queue("sched-model") == [(0, "head")] + + @pytest.fixture(autouse=True) def _vcr_outcome_gate(request, vcr): install_live_call_probe(request, vcr) From 70e6a2ff248893ea3129b625b517f6598b9a3f7b Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 7 Oct 2026 17:39:38 -0700 Subject: [PATCH 03/31] fix(vertex_ai): dial the multi-region Live API host for realtime sessions (#45166) * fix(vertex_ai): dial the multi-region Live API host for realtime sessions A realtime deployment with vertex_location us or eu dialed {location}-aiplatform.googleapis.com, which Vertex does not serve for a multi-region, so the WebSocket handshake came back 404. The URL builder now resolves its host through the shared Vertex host resolver, which already maps us and eu to aiplatform.{geo}.rep.googleapis.com and leaves regions and global unchanged. * test(vertex_ai): annotate the realtime multi-region URL tests * fix(vertex_ai): skip the TLS argument when the realtime api_base is plain ws:// The Vertex realtime session and health-check connects passed the shared TLS context to every websocket dial, so an api_base the transformation already maps from http:// to ws:// failed with 'ssl argument is incompatible with a ws:// URI' before any frame was sent. The OpenAI realtime handler already skips the argument for ws://; the shared helper now does the same for the Vertex path * test(integration): cover the Vertex realtime multi-region host through the proxy Fourteen cells on the scripted Gemini Live upstream: api_base overrides for us, eu, us-central1 and global complete a turn; malformed locations are refused before any dial; a missing location defaults to us-central1; the realtime health check handshakes the override host and reports a malformed location without dialing; a refused handshake reaches the client and the next session connects; twelve concurrent sessions across locations each reach the upstream once; an upstream outage closes every open session and the proxy recovers * test(realtime): type the health-check helpers and flatten the burst order * test(http_handler): type the realtime TLS test signatures --- litellm/llms/custom_httpx/http_handler.py | 12 + litellm/llms/custom_httpx/llm_http_handler.py | 13 +- .../llms/vertex_ai/realtime/transformation.py | 15 +- litellm/realtime_api/main.py | 8 +- tests/integration/_support/upstream.py | 34 ++ .../test_vertex_realtime_multi_region_host.py | 468 ++++++++++++++++++ .../llms/custom_httpx/test_http_handler.py | 41 +- .../test_vertex_ai_realtime_transformation.py | 24 +- tests/unit/realtime_api/test_main.py | 41 ++ 9 files changed, 628 insertions(+), 28 deletions(-) create mode 100644 tests/integration/providers/test_vertex_realtime_multi_region_host.py diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 51f9f1e9fa4..3351d34b16f 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -450,6 +450,18 @@ def get_shared_realtime_ssl_context() -> bool | str | ssl.SSLContext: return _shared_realtime_ssl_context +def realtime_ssl_for_url(url: str) -> bool | str | ssl.SSLContext | None: + if url.startswith("ws://"): + return None + shared: Final = get_shared_realtime_ssl_context() + if shared is not False: + return shared + unverified: Final = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + unverified.check_hostname = False + unverified.verify_mode = ssl.CERT_NONE + return unverified + + def mask_sensitive_info(error_message): # Find the start of the key parameter if isinstance(error_message, str): diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 2b86b9bbe5c..b57ce8ea2e8 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -197,7 +197,7 @@ def _rust_responses_websocket_enabled( return decision(context) is not Decision.PYTHON -from .http_handler import get_shared_realtime_ssl_context +from .http_handler import get_shared_realtime_ssl_context, realtime_ssl_for_url if TYPE_CHECKING: from aiohttp import ClientSession @@ -258,7 +258,7 @@ class _WebsocketsModule(Protocol): *, additional_headers: Mapping[str, str], max_size: int | None, - ssl: bool | str | ssl.SSLContext, + ssl: bool | str | ssl.SSLContext | None, open_timeout: float, ) -> Awaitable["ClientConnection"]: ... @@ -6257,7 +6257,7 @@ class BaseLLMHTTPHandler: websockets_module: _WebsocketsModule, url: str, headers: dict, - ssl_context: bool | str | ssl.SSLContext, + ssl_context: bool | str | ssl.SSLContext | None, *, open_timeout: float = 8.0, max_attempts: int = 3, @@ -6329,12 +6329,7 @@ class BaseLLMHTTPHandler: ) try: - ssl_context = get_shared_realtime_ssl_context() - if url.startswith("wss://") and ssl_context is False: - # Keep TLS for wss:// while honoring SSL_VERIFY=False semantics. - ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) - ssl_context.check_hostname = False - ssl_context.verify_mode = ssl.CERT_NONE + ssl_context: Final = realtime_ssl_for_url(url) provider_backend: Final = await provider_config.open_backend(url, headers) backend_ws: Final = ( provider_backend diff --git a/litellm/llms/vertex_ai/realtime/transformation.py b/litellm/llms/vertex_ai/realtime/transformation.py index 9fed6d52f0e..c14dd535145 100644 --- a/litellm/llms/vertex_ai/realtime/transformation.py +++ b/litellm/llms/vertex_ai/realtime/transformation.py @@ -5,9 +5,14 @@ Extends GeminiRealtimeConfig but adapts the WSS URL and auth header for the Vertex AI endpoint instead of Google AI Studio. URL pattern: - wss://{location}-aiplatform.googleapis.com/ws/ + wss://{vertex host for the location}/ws/ google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent +The host is the one ``litellm.llms.vertex_ai.common_utils.get_vertex_base_url`` +resolves for the location: ``{region}-aiplatform.googleapis.com`` for a region, +``aiplatform.{geo}.rep.googleapis.com`` for the ``us`` / ``eu`` multi-regions, +and ``aiplatform.googleapis.com`` for ``global``. + Auth: OAuth2 Bearer token (not an API key). """ @@ -21,6 +26,7 @@ from litellm.llms.vertex_ai.audio_transcription.realtime_transformation import ( VertexChirpRealtimeConfig, is_vertex_speech_to_text_model, ) +from litellm.llms.vertex_ai.common_utils import get_vertex_base_url from litellm.llms.vertex_ai.vertex_llm_base import VertexBase @@ -63,12 +69,7 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig): base = base.replace("https://", "wss://").replace("http://", "ws://") return f"{base}/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent" - location: Final = self._location - if location == "global": - host = "aiplatform.googleapis.com" - else: - host = f"{location}-aiplatform.googleapis.com" - + host: Final = get_vertex_base_url(self._location).removeprefix("https://") return f"wss://{host}/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent" # ------------------------------------------------------------------ diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 83cacf6c4d9..0e79cbfda13 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -36,7 +36,7 @@ from ..litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from ..llms.azure.common_utils import get_azure_ad_token from ..llms.azure.realtime.handler import AzureOpenAIRealtime, azure_realtime_protocol_for_client from ..llms.bedrock.realtime.handler import BedrockRealtime -from ..llms.custom_httpx.http_handler import get_shared_realtime_ssl_context +from ..llms.custom_httpx.http_handler import realtime_ssl_for_url from ..llms.openai.realtime.handler import OpenAIRealtime from ..llms.vertex_ai.audio_transcription.realtime_transformation import is_vertex_speech_to_text_model from ..llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig, vertex_realtime_config @@ -721,23 +721,21 @@ async def realtime_health_check( location=resolved_location, ) url = vertex_realtime_config.get_complete_url(api_base=resolved_api_base, model=model) - vertex_ssl_context: Final = get_shared_realtime_ssl_context() headers: Final = vertex_realtime_config.validate_environment(headers={}, model=model, api_key=None) async with websockets.connect( url, additional_headers=headers, max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, - ssl=vertex_ssl_context, + ssl=realtime_ssl_for_url(url), ): return True else: raise ValueError(f"Unsupported model: {model}") - ssl_context: Final = get_shared_realtime_ssl_context() async with websockets.connect( url, additional_headers=auth_headers, max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, - ssl=ssl_context, + ssl=realtime_ssl_for_url(url), ): return True diff --git a/tests/integration/_support/upstream.py b/tests/integration/_support/upstream.py index ae87768cd06..0c7dab90fcc 100644 --- a/tests/integration/_support/upstream.py +++ b/tests/integration/_support/upstream.py @@ -440,6 +440,37 @@ class Provider: while pending: await websocket.send_json(_rendered_realtime_event(pending.popleft(), scenario_id)) + async def gemini_live(self, websocket: WebSocket) -> None: + authorization: Final = websocket.headers.get("authorization", "") + scenario_id: Final = websocket.headers.get("x-goog-user-project", "") + self.observations.put( + Observation( + websocket.url.path, + authorization, + {"host": websocket.headers.get("host", "")}, + "WEBSOCKET", + scenario_id, + ) + ) + response: Final = self.scenario_store.get(scenario_id) + if not isinstance(response, RealtimeResponse): + await websocket.close(code=4404) + return + await websocket.accept() + pending: Final = deque(response.events) + async for message in websocket.iter_json(): + payload: Final = JSON_OBJECT.validate_python(message) + self.observations.put( + Observation(websocket.url.path, authorization, payload, "WEBSOCKET_FRAME", scenario_id) + ) + if "setup" in payload: + await websocket.send_json({"setupComplete": {}}) + continue + if not _GEMINI_LIVE_TRIGGERS.intersection(payload): + continue + while pending: + await websocket.send_json(_rendered_realtime_event(pending.popleft(), scenario_id)) + @staticmethod def _response(response: StoredResponse, scenario_id: str) -> Response: unique_id: Final = f"{scenario_id}-{uuid.uuid4().hex[:8]}" @@ -539,6 +570,7 @@ class Provider: WebSocketRoute("/openai/v1/realtime", self.realtime), WebSocketRoute("/openai/realtime", self.realtime), WebSocketRoute("/v1/asr/realtime", self.muse_realtime), + WebSocketRoute(GEMINI_LIVE_PATH, self.gemini_live), ] ) @@ -562,6 +594,8 @@ def _interaction_body(interaction_id: str, state: InteractionState) -> dict[str, _REALTIME_TRIGGERS: Final = frozenset({"response.create", "input_audio_buffer.commit"}) +_GEMINI_LIVE_TRIGGERS: Final = frozenset({"clientContent", "realtimeInput"}) +GEMINI_LIVE_PATH: Final = "/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent" _TRANSCRIPTION_UPDATE_REFUSED: Final = "Passing a realtime session update to a transcription session is not allowed." _REALTIME_UPDATE_REFUSED: Final = "Passing a transcription session update to a realtime session is not allowed." _NESTED_TURN_DETECTION_TYPE: Final = "session.audio.input.turn_detection.type" diff --git a/tests/integration/providers/test_vertex_realtime_multi_region_host.py b/tests/integration/providers/test_vertex_realtime_multi_region_host.py new file mode 100644 index 00000000000..df279aa2ced --- /dev/null +++ b/tests/integration/providers/test_vertex_realtime_multi_region_host.py @@ -0,0 +1,468 @@ +"""Vertex AI Live sessions resolve their host through the shared Vertex resolver. + +``VertexAIRealtimeConfig.get_complete_url`` dials ``aiplatform.{us,eu}.rep.googleapis.com`` for the +multi-regions, ``{region}-aiplatform.googleapis.com`` for a region and ``aiplatform.googleapis.com`` for +``global``, and a malformed location fails the shared validator before any socket opens. Every row here +runs against the scripted upstream on 127.0.0.1: an ``api_base`` override pins the path and the host the +upstream sees, the setup frame it receives names the location the proxy resolved, and the malformed rows +pin the validator's answer, which the proxy gives without dialing anything. +""" + +from __future__ import annotations + +import asyncio +import json +import uuid +from collections.abc import AsyncIterator, Callable +from dataclasses import dataclass +from hashlib import sha256 +from itertools import chain, repeat +from pathlib import Path +from typing import Final + +import httpx +import pytest +import websockets +from pydantic import JsonValue +from websockets.asyncio.client import ClientConnection +from websockets.exceptions import ConnectionClosed + +from tests.integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value, string_value +from tests.integration._support.database import read_rows +from tests.integration._support.process import owned_upstream +from tests.integration._support.upstream import GEMINI_LIVE_PATH, ScenarioHandle, delete_scenario, register_scenario +from tests.integration.cost_calculation.cost_tracking_case import RealtimeResponse + +RecordProperty = Callable[[str, object], None] + +pytestmark: Final = pytest.mark.timeout(180) + +LIVE_MODEL: Final = "gemini-3.8-live" +LOCATIONS: Final = ("us", "eu", "us-central1", "global") +MALFORMED_LOCATIONS: Final = ("US", "us/evil") +DEFAULT_LOCATION: Final = "us-central1" +INVALID_LOCATION: Final = "Invalid vertex_location format" +HANDSHAKE_REFUSED: Final = "Upstream realtime handshake rejected with HTTP 403" +INTERNAL_CLOSE: Final = 1011 +REFUSAL_CLOSE: Final = 1008 +SESSIONS_PER_LOCATION: Final = 3 +OUTAGE_BURST: Final = 20 +INPUT_TOKENS: Final = 7 +OUTPUT_TOKENS: Final = 5 +CONVERGENCE_SECONDS: Final = 30 +SPEND_SQL: Final = 'SELECT call_type, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE api_key = %s' + + +@dataclass(frozen=True, slots=True) +class Session: + events: tuple[dict[str, JsonValue], ...] + + @property + def types(self) -> tuple[str, ...]: + return tuple(string_value(event["type"]) for event in self.events) + + @property + def close_code(self) -> int | None: + closes: Final = tuple(event for event in self.events if event["type"] == "closed") + return _integer(closes[-1]["code"]) if closes else None + + @property + def session_model(self) -> str: + return string_value(object_value(self.events[0]["session"])["model"]) + + @property + def response_ids(self) -> tuple[str, ...]: + done: Final = tuple(event for event in self.events if event["type"] == "response.done") + return tuple(string_value(object_value(event["response"])["id"]) for event in done) + + @property + def error_messages(self) -> tuple[str, ...]: + errors: Final = tuple(event for event in self.events if event["type"] == "error") + return tuple(string_value(object_value(event["error"])["message"]) for event in errors) + + @property + def text(self) -> str: + deltas: Final = tuple(event for event in self.events if event["type"] == "response.output_text.delta") + return "".join(string_value(event["delta"]) for event in deltas) + + def completed_turn(self, scenario_id: str) -> bool: + return ( + self.types[0] == "session.created" + and self.types[-1] == "response.done" + and self.session_model == LIVE_MODEL + and self.text == f"scripted {scenario_id}" + and len(self.response_ids) == 1 + ) + + +def _integer(value: JsonValue) -> int: + assert isinstance(value, int), value + return value + + +def _ws_base(http_url: str) -> str: + return http_url.replace("https://", "wss://").replace("http://", "ws://") + + +def _proxy_url(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + +def _upstream_url(gateway: Gateway) -> str: + return gateway.upstream_url.rstrip("/") + + +def _scripted_turn() -> RealtimeResponse: + return RealtimeResponse( + content_type="application/x-realtime", + events=( + {"serverContent": {"modelTurn": {"parts": [{"text": "scripted $REQUEST_ID"}]}}}, + { + "serverContent": {"turnComplete": True}, + "usageMetadata": { + "promptTokenCount": INPUT_TOKENS, + "responseTokenCount": OUTPUT_TOKENS, + "totalTokenCount": INPUT_TOKENS + OUTPUT_TOKENS, + }, + }, + ), + ) + + +def _scripted(scenario: Scenario) -> ScenarioHandle: + handle: Final = register_scenario(f"vertex-live-{uuid.uuid4().hex[:12]}", _scripted_turn()) + scenario.cleanups.callback(delete_scenario, handle) + return handle + + +def _credentials(upstream_url: str) -> str: + return json.dumps( + { + "type": "external_account", + "audience": "synthetic-vertex-audience", + "subject_token_type": "urn:ietf:params:oauth:token-type:jwt", + "token_url": f"{upstream_url}/_oauth/token", + "credential_source": {"url": f"{upstream_url}/health"}, + } + ) + + +def _deployment( + gateway: Gateway, scenario: Scenario, project: str, *, location: str | None, api_base: str | None +) -> str: + name: Final = f"vertex-live-{uuid.uuid4().hex[:12]}" + litellm_params: Final[dict[str, JsonValue]] = { + "model": f"vertex_ai/{LIVE_MODEL}", + "vertex_project": project, + "vertex_credentials": _credentials(_upstream_url(gateway)), + **({} if location is None else {"vertex_location": location}), + **({} if api_base is None else {"api_base": api_base}), + } + created: Final = gateway.post( + "/model/new", {"model_name": name, "litellm_params": litellm_params, "model_info": {"mode": "realtime"}} + ) + scenario.cleanups.callback(scenario.delete_model, string_value(object_value(created["model_info"])["id"])) + return name + + +def _model_path(project: str, location: str) -> str: + return f"projects/{project}/locations/{location}/publishers/google/models/{LIVE_MODEL}" + + +def _user_turn() -> str: + return json.dumps( + { + "type": "conversation.item.create", + "item": { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": "say the scripted line"}], + }, + } + ) + + +def _close_code(closed: ConnectionClosed) -> int: + return 1006 if closed.rcvd is None else closed.rcvd.code + + +async def _frames(socket: ClientConnection, turns: int) -> AsyncIterator[dict[str, JsonValue]]: + try: + first: Final = JSON_OBJECT.validate_json(await socket.recv()) + yield first + if first.get("type") != "session.created": + async for message in socket: + yield JSON_OBJECT.validate_json(message) + return + for _ in range(turns): + await socket.send(_user_turn()) + async for message in socket: + event: Final = JSON_OBJECT.validate_json(message) + yield event + if event.get("type") == "response.done": + break + except ConnectionClosed as closed: + yield {"type": "closed", "code": _close_code(closed)} + + +async def _collect(socket: ClientConnection, turns: int) -> tuple[dict[str, JsonValue], ...]: + return tuple([frame async for frame in _frames(socket, turns)]) + + +async def _session(ws_base: str, model: str, key: str, *, turns: int = 1) -> Session: + headers: Final = {"Authorization": f"Bearer {key}"} + async with websockets.connect(f"{ws_base}/v1/realtime?model={model}", additional_headers=headers) as socket: + return Session(await asyncio.wait_for(_collect(socket, turns), 60)) + + +def _run(gateway: Gateway, model: str, key: str, *, turns: int = 1) -> Session: + return asyncio.run(_session(_ws_base(_proxy_url(gateway)), model, key, turns=turns)) + + +def _observations(gateway: Gateway) -> tuple[dict[str, JsonValue], ...]: + with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream: + return tuple(map(object_value, upstream.get("/__observations").json()["requests"])) + + +def _for_project(observed: tuple[dict[str, JsonValue], ...], project: str) -> tuple[dict[str, JsonValue], ...]: + return tuple(request for request in observed if request["api_key"] == project) + + +def _upgrade_hosts(observed: tuple[dict[str, JsonValue], ...]) -> tuple[str, ...]: + upgrades: Final = tuple(request for request in observed if request["method"] == "WEBSOCKET") + assert all(request["path"] == GEMINI_LIVE_PATH for request in upgrades), upgrades + return tuple(string_value(object_value(request["body"])["host"]) for request in upgrades) + + +def _setup_models(observed: tuple[dict[str, JsonValue], ...]) -> tuple[str, ...]: + frames: Final = tuple( + object_value(request["body"]) for request in observed if request["method"] == "WEBSOCKET_FRAME" + ) + return tuple(string_value(object_value(frame["setup"])["model"]) for frame in frames if "setup" in frame) + + +def _authority(gateway: Gateway) -> str: + return _upstream_url(gateway).removeprefix("http://") + + +def _spend_rows(key: str, count: int) -> list[dict[str, JsonValue]]: + return eventually( + lambda: read_rows(SPEND_SQL, (sha256(key.encode()).hexdigest(),)), + lambda rows: len(rows) == count, + seconds=70, + ) + + +def _billed(rows: list[dict[str, JsonValue]]) -> list[tuple[JsonValue, JsonValue, JsonValue]]: + return [(row["call_type"], row["prompt_tokens"], row["completion_tokens"]) for row in rows] + + +def _health(gateway: Gateway, model: str) -> httpx.Response: + return gateway.request("GET", "/health", params={"model": model}) + + +def _probed_count(response: httpx.Response) -> int: + body: Final = JSON_OBJECT.validate_json(response.content) + return _integer(body.get("healthy_count", 0)) + _integer(body.get("unhealthy_count", 0)) + + +def _converged_health(gateway: Gateway, model: str) -> httpx.Response: + return eventually( + lambda: _health(gateway, model), lambda response: _probed_count(response) == 1, seconds=CONVERGENCE_SECONDS + ) + + +def _assert_scripted_turn_reached_upstream(gateway: Gateway, project: str, location: str, session: Session) -> None: + assert session.completed_turn(project), session + observed: Final = _for_project(_observations(gateway), project) + assert _upgrade_hosts(observed) == (_authority(gateway),), observed + assert _setup_models(observed) == (_model_path(project, location),), observed + + +@pytest.mark.parametrize("location", LOCATIONS) +def test_api_base_override_completes_a_turn_for_every_location(gateway: Gateway, location: str) -> None: + with gateway.scenario() as scenario: + handle: Final = _scripted(scenario) + key: Final = scenario.key() + model: Final = _deployment( + gateway, scenario, handle.scenario_id, location=location, api_base=_upstream_url(gateway) + ) + session: Final = _run(gateway, model, key) + _assert_scripted_turn_reached_upstream(gateway, handle.scenario_id, location, session) + assert _billed(_spend_rows(key, 1)) == [("_arealtime", INPUT_TOKENS, OUTPUT_TOKENS)] + + +@pytest.mark.parametrize("location", MALFORMED_LOCATIONS) +def test_malformed_location_is_refused_by_the_validator_before_any_dial(gateway: Gateway, location: str) -> None: + with gateway.scenario() as scenario: + handle: Final = _scripted(scenario) + key: Final = scenario.key() + malformed: Final = _deployment(gateway, scenario, handle.scenario_id, location=location, api_base=None) + refused: Final = _run(gateway, malformed, key) + assert refused.types == ("error", "closed"), refused + assert refused.error_messages == (INVALID_LOCATION,), refused + assert refused.close_code == INTERNAL_CLOSE, refused + model: Final = _deployment( + gateway, scenario, handle.scenario_id, location="us", api_base=_upstream_url(gateway) + ) + session: Final = _run(gateway, model, key) + _assert_scripted_turn_reached_upstream(gateway, handle.scenario_id, "us", session) + + +@pytest.mark.parametrize("location", [None, ""], ids=["omitted", "empty"]) +def test_missing_location_defaults_to_us_central1(gateway: Gateway, location: str | None) -> None: + with gateway.scenario() as scenario: + handle: Final = _scripted(scenario) + key: Final = scenario.key() + model: Final = _deployment( + gateway, scenario, handle.scenario_id, location=location, api_base=_upstream_url(gateway) + ) + session: Final = _run(gateway, model, key) + _assert_scripted_turn_reached_upstream(gateway, handle.scenario_id, DEFAULT_LOCATION, session) + + +@pytest.mark.parametrize("location", ["us", "eu"]) +def test_health_check_handshakes_with_the_override_host(gateway: Gateway, location: str) -> None: + with gateway.scenario() as scenario: + handle: Final = _scripted(scenario) + model: Final = _deployment( + gateway, scenario, handle.scenario_id, location=location, api_base=_upstream_url(gateway) + ) + health: Final = _converged_health(gateway, model) + assert health.status_code == 200, health.text + body: Final = JSON_OBJECT.validate_json(health.content) + assert (body["healthy_count"], body["unhealthy_count"]) == (1, 0), health.text + observed: Final = _for_project(_observations(gateway), handle.scenario_id) + assert _upgrade_hosts(observed) == (_authority(gateway),), observed + assert _setup_models(observed) == (), observed + + +def test_health_check_reports_a_malformed_location_without_dialing(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + handle: Final = _scripted(scenario) + model: Final = _deployment(gateway, scenario, handle.scenario_id, location="US", api_base=None) + health: Final = _converged_health(gateway, model) + assert health.status_code == 503, health.text + body: Final = JSON_OBJECT.validate_json(health.content) + assert (body["healthy_count"], body["unhealthy_count"]) == (0, 1), health.text + unhealthy: Final = body["unhealthy_endpoints"] + assert isinstance(unhealthy, list) and len(unhealthy) == 1, health.text + assert INVALID_LOCATION in string_value(object_value(unhealthy[0])["error"]), health.text + assert _for_project(_observations(gateway), handle.scenario_id) == (), "the upstream saw a dial" + + +def test_upstream_handshake_refusal_reaches_the_client_and_the_next_session_connects(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + key: Final = scenario.key() + unregistered: Final = f"vertex-live-unknown-{uuid.uuid4().hex[:12]}" + refused_model: Final = _deployment( + gateway, scenario, unregistered, location="us", api_base=_upstream_url(gateway) + ) + refused: Final = _run(gateway, refused_model, key) + assert refused.types == ("error", "closed"), refused + assert refused.error_messages == (HANDSHAKE_REFUSED,), refused + assert refused.close_code == REFUSAL_CLOSE, refused + assert _upgrade_hosts(_for_project(_observations(gateway), unregistered)) == (_authority(gateway),) + handle: Final = _scripted(scenario) + model: Final = _deployment( + gateway, scenario, handle.scenario_id, location="us", api_base=_upstream_url(gateway) + ) + session: Final = _run(gateway, model, key) + _assert_scripted_turn_reached_upstream(gateway, handle.scenario_id, "us", session) + chat: Final = gateway.chat(scenario.model(), key=key) + assert string_value(chat["object"]) == "chat.completion", chat + + +async def _burst(ws_base: str, models: tuple[str, ...], key: str) -> tuple[Session, ...]: + return tuple(await asyncio.gather(*(_session(ws_base, model, key) for model in models))) + + +def test_concurrent_sessions_across_locations_each_reach_the_upstream_once(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + key: Final = scenario.key() + handles: Final = {location: _scripted(scenario) for location in LOCATIONS} + models: Final = { + location: _deployment( + gateway, scenario, handles[location].scenario_id, location=location, api_base=_upstream_url(gateway) + ) + for location in LOCATIONS + } + order: Final = tuple(chain.from_iterable(repeat(location, SESSIONS_PER_LOCATION) for location in LOCATIONS)) + sessions: Final = asyncio.run( + _burst(_ws_base(_proxy_url(gateway)), tuple(models[location] for location in order), key) + ) + assert all( + session.completed_turn(handles[location].scenario_id) for location, session in zip(order, sessions) + ), sessions + assert len({session.response_ids[0] for session in sessions}) == len(order), sessions + observed: Final = _observations(gateway) + for location, handle in handles.items(): + mine: Final = _for_project(observed, handle.scenario_id) + assert _upgrade_hosts(mine) == (_authority(gateway),) * SESSIONS_PER_LOCATION, mine + assert _setup_models(mine) == (_model_path(handle.scenario_id, location),) * SESSIONS_PER_LOCATION, mine + assert _billed(_spend_rows(key, len(order))) == [("_arealtime", INPUT_TOKENS, OUTPUT_TOKENS)] * len(order) + + +async def _frames_until_closed(socket: ClientConnection) -> AsyncIterator[dict[str, JsonValue]]: + try: + async for message in socket: + yield JSON_OBJECT.validate_json(message) + except ConnectionClosed as closed: + yield {"type": "closed", "code": _close_code(closed)} + + +async def _hold_until_closed(ws_base: str, model: str, key: str, opened: asyncio.Queue[str]) -> Session: + headers: Final = {"Authorization": f"Bearer {key}"} + async with websockets.connect(f"{ws_base}/v1/realtime?model={model}", additional_headers=headers) as socket: + created: Final = JSON_OBJECT.validate_json(await socket.recv()) + assert created.get("type") == "session.created", created + await opened.put(string_value(object_value(created["session"])["id"])) + return Session(tuple([frame async for frame in _frames_until_closed(socket)])) + + +def _relays_the_upstream_close(session: Session) -> bool: + return f"upstream websocket closed with code {session.close_code}" in session.error_messages[0] + + +async def _drain(opened: asyncio.Queue[str], count: int) -> tuple[str, ...]: + return tuple([await opened.get() for _ in range(count)]) + + +async def _burst_through_outage( + ws_base: str, proxy_url: str, model: str, key: str, stop_upstream: Callable[[], None] +) -> tuple[Session, ...]: + opened: Final[asyncio.Queue[str]] = asyncio.Queue() + holders: Final = tuple( + asyncio.ensure_future(_hold_until_closed(ws_base, model, key, opened)) for _ in range(OUTAGE_BURST) + ) + opened_sessions: Final = await asyncio.wait_for(_drain(opened, OUTAGE_BURST), 60) + assert len(opened_sessions) == OUTAGE_BURST, opened_sessions + await asyncio.to_thread(stop_upstream) + async with httpx.AsyncClient(base_url=proxy_url, timeout=15, trust_env=False) as client: + liveliness: Final = await client.get("/health/liveliness") + assert liveliness.status_code == 200, liveliness.text + return tuple(await asyncio.wait_for(asyncio.gather(*holders), 90)) + + +@pytest.mark.timeout(240) +def test_upstream_outage_closes_every_open_session_and_recovers( + gateway: Gateway, tmp_path: Path, record_property: RecordProperty +) -> None: + with gateway.scenario() as scenario, owned_upstream(tmp_path) as slot: + project: Final = f"vertex-live-outage-{uuid.uuid4().hex[:12]}" + register_scenario(project, _scripted_turn(), control_url=slot.url) + key: Final = scenario.key() + model: Final = _deployment(gateway, scenario, project, location="us", api_base=slot.url) + held: Final = asyncio.run( + _burst_through_outage(_ws_base(_proxy_url(gateway)), _proxy_url(gateway), model, key, slot.stop) + ) + record_property("close_codes_during_upstream_outage", sorted(session.close_code or 0 for session in held)) + assert [session.types for session in held] == [("error", "closed")] * OUTAGE_BURST, held + assert all(_relays_the_upstream_close(session) for session in held), held + assert len({session.close_code for session in held}) == 1, held + slot.start() + register_scenario(project, _scripted_turn(), control_url=slot.url) + recovered: Final = _run(gateway, model, key) + assert recovered.completed_turn(project), recovered + rows: Final = _spend_rows(key, OUTAGE_BURST + 1) + assert {str(row["call_type"]) for row in rows} == {"_arealtime"}, rows diff --git a/tests/unit/llms/custom_httpx/test_http_handler.py b/tests/unit/llms/custom_httpx/test_http_handler.py index f5ae65e9898..b6b0db69162 100644 --- a/tests/unit/llms/custom_httpx/test_http_handler.py +++ b/tests/unit/llms/custom_httpx/test_http_handler.py @@ -24,7 +24,9 @@ from litellm.llms.custom_httpx.http_handler import ( HTTPHandler, MaskedHTTPStatusError, get_httpx_client, + get_shared_realtime_ssl_context, get_ssl_configuration, + realtime_ssl_for_url, ) from litellm.types.llms.custom_http import VerifyTypes from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer @@ -1493,7 +1495,6 @@ async def test_connection_error_retry_forwards_content(method: str): await handler.close() - @pytest.fixture def forward_proxy_server(): """Plain HTTP forward proxy that records the absolute URIs it is asked to fetch.""" @@ -1624,9 +1625,7 @@ def private_ca_tls_upstream(tmp_path: pathlib.Path): ca_pem.write_bytes(cert.public_bytes(serialization.Encoding.PEM)) key_pem = tmp_path / "key.pem" key_pem.write_bytes( - key.private_bytes( - serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption() - ) + key.private_bytes(serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption()) ) class OkTlsHandler(BaseHTTPRequestHandler): @@ -1801,8 +1800,11 @@ async def test_bounded_get_preserves_sdk_redirect_auth_and_query_handling(respx_ handler = AsyncHTTPHandler() try: response = await handler.get( - "https://example.com/spec.json?original=1", max_response_bytes=100, follow_redirects=True, - headers={"Authorization": "Bearer sentinel", "Accept-Encoding": "gzip"}, timeout=2.0, + "https://example.com/spec.json?original=1", + max_response_bytes=100, + follow_redirects=True, + headers={"Authorization": "Bearer sentinel", "Accept-Encoding": "gzip"}, + timeout=2.0, ) finally: await handler.close() @@ -1888,6 +1890,7 @@ def _vcr_outcome_gate(request, vcr): yield record_vcr_outcome(request, vcr) + @pytest.fixture(scope="function") def isolate_litellm_state(): """ @@ -1940,6 +1943,7 @@ def isolate_litellm_state(): setattr(litellm, attr, original_value) _invalidate_model_cost_lowercase_map() + _SCALAR_DEFAULTS = { "num_retries": getattr(litellm, "num_retries", None), "num_retries_per_request": getattr(litellm, "num_retries_per_request", None), @@ -1960,6 +1964,7 @@ _SCALAR_DEFAULTS = { "api_key": getattr(litellm, "api_key", None), } + @pytest.fixture(scope="module") def setup_and_teardown(): """ @@ -1982,12 +1987,14 @@ def setup_and_teardown(): litellm.in_memory_llm_clients_cache.flush_cache() yield + _SERVER_DELAY_S = 5 _PER_REQUEST_TIMEOUT_S = 1.0 _CLIENT_DEFAULT_TIMEOUT_S = 60.0 + class _SlowHandler(BaseHTTPRequestHandler): def do_POST(self): time.sleep(_SERVER_DELAY_S) @@ -2001,6 +2008,7 @@ class _SlowHandler(BaseHTTPRequestHandler): def log_message(self, *args): pass + @pytest.mark.usefixtures("_vcr_outcome_gate", "isolate_litellm_state", "setup_and_teardown") def test_post_delay_exceeds_per_request_timeout_raises(): server = ThreadingHTTPServer(("127.0.0.1", 0), _SlowHandler) @@ -2022,3 +2030,24 @@ def test_post_delay_exceeds_per_request_timeout_raises(): handler.close() server.shutdown() server.server_close() + + +def test_realtime_ssl_for_url_sends_no_tls_argument_for_a_plain_ws_endpoint() -> None: + assert ( + realtime_ssl_for_url("ws://127.0.0.1:8080/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent") + is None + ) + + +def test_realtime_ssl_for_url_keeps_the_shared_context_for_wss_endpoints() -> None: + shared: Final = get_shared_realtime_ssl_context() + assert isinstance(shared, ssl.SSLContext) + assert realtime_ssl_for_url("wss://aiplatform.us.rep.googleapis.com/ws") is shared + + +def test_realtime_ssl_for_url_turns_ssl_verify_false_into_an_unverified_tls_context(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr("litellm.llms.custom_httpx.http_handler._shared_realtime_ssl_context", False) + selected: Final = realtime_ssl_for_url("wss://aiplatform.us.rep.googleapis.com/ws") + assert isinstance(selected, ssl.SSLContext) + assert selected.verify_mode == ssl.CERT_NONE + assert selected.check_hostname is False diff --git a/tests/unit/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py b/tests/unit/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py index 14b3bdb48a1..d6986a2697a 100644 --- a/tests/unit/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py +++ b/tests/unit/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py @@ -2,7 +2,7 @@ Unit tests for VertexAIRealtimeConfig. Validates: -- URL construction (regional and global) +- URL construction (regional, multi-region and global) - Auth headers (Bearer token + project header) - Session setup message format - Full text-in / text-out round-trip via RealTimeStreaming with a mocked @@ -10,12 +10,14 @@ Validates: """ import json +from typing import Final from unittest.mock import AsyncMock, MagicMock import pytest import websockets.exceptions # registers websockets.exceptions on the websockets namespace import litellm +from litellm.llms.vertex_ai.common_utils import get_vertex_base_url from litellm.llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig # --------------------------------------------------------------------------- @@ -45,6 +47,26 @@ def test_get_complete_url_global(): ) +@pytest.mark.parametrize("location", ["us", "eu"]) +def test_get_complete_url_multi_region_uses_rep_host(location: str): + cfg: Final = VertexAIRealtimeConfig(access_token="tok", project="my-proj", location=location) + url: Final = cfg.get_complete_url(api_base=None, model="gemini-3.8-live") + # Google documents the multi-region Vertex endpoints as aiplatform.{us,eu}.rep.googleapis.com + # (https://docs.cloud.google.com/vertex-ai/generative-ai/docs/learn/locations, read 2026-10-07) + assert url == ( + f"wss://aiplatform.{location}.rep.googleapis.com" + "/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent" + ) + + +@pytest.mark.parametrize("location", ["us", "eu", "global", "us-central1", "europe-west4"]) +def test_get_complete_url_host_matches_shared_vertex_host(location: str): + cfg: Final = VertexAIRealtimeConfig(access_token="tok", project="my-proj", location=location) + url: Final = cfg.get_complete_url(api_base=None, model="gemini-3.8-live") + shared_host: Final = get_vertex_base_url(location).removeprefix("https://") + assert url == f"wss://{shared_host}/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent" + + def test_get_complete_url_custom_api_base(): cfg = VertexAIRealtimeConfig( access_token="tok", project="my-proj", location="us-central1" diff --git a/tests/unit/realtime_api/test_main.py b/tests/unit/realtime_api/test_main.py index 25f116f5991..a692f38ec04 100644 --- a/tests/unit/realtime_api/test_main.py +++ b/tests/unit/realtime_api/test_main.py @@ -658,6 +658,47 @@ async def test_arealtime_drops_model_from_the_upstream_url_only_for_transcriptio assert connect.url == expected_backend_url +async def _vertex_health_check_connect_for(monkeypatch: pytest.MonkeyPatch, api_base: str | None) -> _CapturingConnect: + async def fake_token_resolver( + credentials: object, project_id: str | None, custom_llm_provider: str + ) -> tuple[str, str]: + return "access-token", project_id or "" + + monkeypatch.setattr(realtime_main, "vertex_access_token_resolver", fake_token_resolver) + connect: Final = _CapturingConnect() + with patch("websockets.connect", connect): + assert await realtime_main._realtime_health_check( + model="gemini-3.8-live", + custom_llm_provider="vertex_ai", + api_key=None, + api_base=api_base, + model_params={ + "vertex_project": "proj-1", + "vertex_credentials": "fake-credentials", + "vertex_location": "us", + }, + ) + return connect + + +@pytest.mark.asyncio +async def test_vertex_health_check_sends_no_tls_argument_to_a_plain_ws_api_base( + monkeypatch: pytest.MonkeyPatch, +) -> None: + connect: Final = await _vertex_health_check_connect_for(monkeypatch, "http://127.0.0.1:8080") + assert connect.url == "ws://127.0.0.1:8080/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent" + assert connect.kwargs["ssl"] is None + + +@pytest.mark.asyncio +async def test_vertex_health_check_keeps_tls_for_the_multi_region_host(monkeypatch: pytest.MonkeyPatch) -> None: + connect: Final = await _vertex_health_check_connect_for(monkeypatch, None) + assert connect.url == ( + "wss://aiplatform.us.rep.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent" + ) + assert connect.kwargs["ssl"] is not None + + BLOCKED_PHRASE = "XSECRETBLOCKTESTPHRASEX" From 5531169710aeb132708f90aa8e4bd36b7b467a25 Mon Sep 17 00:00:00 2001 From: ishaan-berri <155045088+ishaan-berri@users.noreply.github.com> Date: Wed, 7 Oct 2026 17:41:13 -0700 Subject: [PATCH 04/31] feat(ui): add Moyai cloud coding agent to the view switcher (#45196) * feat(ui): add Moyai cloud coding agent to the view switcher Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> * fix(ui): restore hover scale on Moyai GitHub CTA Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_types.py | 1 + litellm/proxy/moyai_endpoints.py | 260 ++++++++++ litellm/proxy/proxy_server.py | 2 + .../proxy_setting_endpoints.py | 27 +- tests/unit/proxy/test_moyai_endpoints.py | 290 +++++++++++ .../test_proxy_setting_endpoints.py | 68 +++ .../public/assets/moyai/logos/anthropic.svg | 5 + .../public/assets/moyai/logos/bedrock.svg | 1 + .../public/assets/moyai/logos/deepseek.svg | 25 + .../public/assets/moyai/logos/fireworks.svg | 1 + .../public/assets/moyai/logos/google.svg | 2 + .../public/assets/moyai/logos/hermes.png | Bin 0 -> 1100 bytes .../public/assets/moyai/logos/langchain.svg | 1 + .../public/assets/moyai/logos/mistral.svg | 1 + .../public/assets/moyai/logos/openai.svg | 5 + .../public/assets/moyai/logos/opencode.svg | 16 + .../public/assets/moyai/logos/xai.svg | 28 ++ .../public/assets/moyai/moyai-head.svg | 19 + .../(dashboard)/hooks/useDisableBlogPosts.ts | 2 +- .../hooks/useDisableBouncingIcon.ts | 2 +- .../hooks/useDisableShowNewBadge.ts | 2 +- .../hooks/useDisableShowPrompts.ts | 2 +- .../hooks/useHideAutoRouterAnnouncement.ts | 2 +- ui/litellm-dashboard/src/app/moyai/page.tsx | 99 ++++ .../components/Navbar/ViewSwitcher.test.tsx | 81 ++- .../src/components/Navbar/ViewSwitcher.tsx | 42 +- .../UISettings/UISettings.test.tsx | 39 ++ .../AdminSettings/UISettings/UISettings.tsx | 36 ++ .../components/moyai/MoyaiConnected.test.tsx | 28 ++ .../src/components/moyai/MoyaiConnected.tsx | 84 ++++ .../components/moyai/MoyaiLanding.module.css | 84 ++++ .../components/moyai/MoyaiLanding.test.tsx | 56 +++ .../src/components/moyai/MoyaiLanding.tsx | 463 ++++++++++++++++++ .../src/components/moyai/moyaiConnect.test.ts | 24 + .../src/components/moyai/moyaiConnect.ts | 14 + .../src/components/moyai/moyaiSky.ts | 168 +++++++ .../src/components/networking.tsx | 7 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 128 +++++ 38 files changed, 2101 insertions(+), 14 deletions(-) create mode 100644 litellm/proxy/moyai_endpoints.py create mode 100644 tests/unit/proxy/test_moyai_endpoints.py create mode 100644 ui/litellm-dashboard/public/assets/moyai/logos/anthropic.svg create mode 100644 ui/litellm-dashboard/public/assets/moyai/logos/bedrock.svg create mode 100644 ui/litellm-dashboard/public/assets/moyai/logos/deepseek.svg create mode 100644 ui/litellm-dashboard/public/assets/moyai/logos/fireworks.svg create mode 100644 ui/litellm-dashboard/public/assets/moyai/logos/google.svg create mode 100644 ui/litellm-dashboard/public/assets/moyai/logos/hermes.png create mode 100644 ui/litellm-dashboard/public/assets/moyai/logos/langchain.svg create mode 100644 ui/litellm-dashboard/public/assets/moyai/logos/mistral.svg create mode 100644 ui/litellm-dashboard/public/assets/moyai/logos/openai.svg create mode 100644 ui/litellm-dashboard/public/assets/moyai/logos/opencode.svg create mode 100644 ui/litellm-dashboard/public/assets/moyai/logos/xai.svg create mode 100644 ui/litellm-dashboard/public/assets/moyai/moyai-head.svg create mode 100644 ui/litellm-dashboard/src/app/moyai/page.tsx create mode 100644 ui/litellm-dashboard/src/components/moyai/MoyaiConnected.test.tsx create mode 100644 ui/litellm-dashboard/src/components/moyai/MoyaiConnected.tsx create mode 100644 ui/litellm-dashboard/src/components/moyai/MoyaiLanding.module.css create mode 100644 ui/litellm-dashboard/src/components/moyai/MoyaiLanding.test.tsx create mode 100644 ui/litellm-dashboard/src/components/moyai/MoyaiLanding.tsx create mode 100644 ui/litellm-dashboard/src/components/moyai/moyaiConnect.test.ts create mode 100644 ui/litellm-dashboard/src/components/moyai/moyaiConnect.ts create mode 100644 ui/litellm-dashboard/src/components/moyai/moyaiSky.ts diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 0ed2b38cfa6..7f92573fde7 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -853,6 +853,7 @@ class LiteLLMRoutes(enum.Enum): "/public/mcp_hub", "/public/skill_hub", "/public/litellm_model_cost_map", + "/moyai/connect/exchange", ) ) diff --git a/litellm/proxy/moyai_endpoints.py b/litellm/proxy/moyai_endpoints.py new file mode 100644 index 00000000000..655fc0560f4 --- /dev/null +++ b/litellm/proxy/moyai_endpoints.py @@ -0,0 +1,260 @@ +"""Moyai quick-connect endpoints. + +`/moyai/connect/start` hands a proxy admin a signed, single-use code pointing +at their Moyai deployment. `/moyai/connect/exchange` trades that code for a +fresh virtual key and persists the deployment as the `moyai_url` UI setting. +The signed code is the credential for the exchange, so it must stay short +lived and single use. +""" + +import base64 +import hashlib +import hmac +import json +import os +import secrets +import time +from typing import TYPE_CHECKING, Annotated, Final +from urllib.parse import urlencode, urlparse + +from fastapi import APIRouter, Depends, HTTPException, Request, status +from pydantic import BaseModel + +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import ( + UI_SETTINGS_CACHE_KEY, + UI_SETTINGS_CACHE_TTL, + _ui_settings_db, + normalize_moyai_url, +) +from litellm.repositories.config_repository import ConfigRepository +from litellm.repositories.table_repositories import UISettingsRepository + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + +router: Final = APIRouter() + +_MOYAI_CODE_TTL_SECONDS: Final = 600 +_MOYAI_NONCE_CONFIG_PREFIX: Final = "moyai_connect_nonce:" +_MOYAI_CONNECT_EXCHANGE_ROUTE: Final = "/moyai/connect/exchange" + + +class MoyaiConnectStartRequest(BaseModel): + moyai_url: str + return_to: str + + +class MoyaiConnectStartResponse(BaseModel): + connect_url: str + + +class MoyaiConnectExchangeRequest(BaseModel): + code: str + moyai_url: str + + +class MoyaiConnectExchangeResponse(BaseModel): + api_key: str + key_alias: str + api_base: str + + +def _b64url(data: bytes) -> str: + return base64.urlsafe_b64encode(data).decode().rstrip("=") + + +def _b64url_decode(data: str) -> bytes: + return base64.urlsafe_b64decode(data + "=" * (-len(data) % 4)) + + +def _origin(url: str) -> str: + parsed: Final = urlparse(url) + return f"{parsed.scheme}://{parsed.netloc}" + + +def _master_key_hmac_key(master_key: str) -> bytes: + return hashlib.sha256(master_key.encode()).digest() + + +def _gateway_url(request: Request) -> str: + if os.environ.get("PROXY_BASE_URL"): + return os.environ["PROXY_BASE_URL"].rstrip("/") + return str(request.base_url).rstrip("/") + + +_MOYAI_KEY_ALLOWED_ROUTES: Final = ["openai_routes", "anthropic_routes", "/model/info"] + + +def _sign_connect_code(master_key: str, moyai_url: str, user_id: str | None) -> str: + payload: Final = json.dumps( + { + "moyai_origin": _origin(moyai_url), + "user_id": user_id, + "exp": int(time.time()) + _MOYAI_CODE_TTL_SECONDS, + "nonce": secrets.token_urlsafe(16), + }, + separators=(",", ":"), + sort_keys=True, + ).encode() + signature: Final = hmac.new(_master_key_hmac_key(master_key), payload, hashlib.sha256).digest() + return f"{_b64url(payload)}.{_b64url(signature)}" + + +def _decode_connect_code(master_key: str, code: str) -> dict: + try: + payload_b64, signature_b64 = code.split(".", 1) + payload_raw: Final = _b64url_decode(payload_b64) + signature: Final = _b64url_decode(signature_b64) + except (ValueError, TypeError): + raise HTTPException(status_code=400, detail="Invalid Moyai connect code") + expected: Final = hmac.new(_master_key_hmac_key(master_key), payload_raw, hashlib.sha256).digest() + if not hmac.compare_digest(signature, expected): + raise HTTPException(status_code=400, detail="Invalid Moyai connect code") + try: + payload: Final = json.loads(payload_raw) + except (ValueError, TypeError): + raise HTTPException(status_code=400, detail="Invalid Moyai connect code") + if not isinstance(payload, dict) or not isinstance(payload.get("exp"), int) or payload["exp"] < int(time.time()): + raise HTTPException(status_code=400, detail="Invalid Moyai connect code") + if not isinstance(payload.get("moyai_origin"), str) or not isinstance(payload.get("nonce"), str): + raise HTTPException(status_code=400, detail="Invalid Moyai connect code") + return payload + + +@router.post( + "/moyai/connect/start", + response_model=MoyaiConnectStartResponse, + tags=["moyai"], +) +async def moyai_connect_start( + request: Request, + body: MoyaiConnectStartRequest, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +) -> MoyaiConnectStartResponse: + from litellm.proxy.proxy_server import master_key + + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value: + raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Only proxy admins can connect Moyai") + + try: + moyai_url: Final = normalize_moyai_url(body.moyai_url) + except ValueError as e: + raise HTTPException(status_code=400, detail=str(e)) + if moyai_url is None: + raise HTTPException(status_code=400, detail="moyai_url is required") + + return_to_parsed: Final = urlparse(body.return_to) + if return_to_parsed.scheme not in ("http", "https") or not return_to_parsed.netloc: + raise HTTPException(status_code=400, detail="return_to must be an absolute http or https URL") + + if not master_key: + raise HTTPException( + status_code=400, + detail="Moyai quick connect needs LITELLM_MASTER_KEY set on the proxy", + ) + + code: Final = _sign_connect_code(master_key, moyai_url, user_api_key_dict.user_id) + connect_url: Final = f"{moyai_url}/connect/litellm?" + urlencode( + {"gateway_url": _gateway_url(request), "code": code, "return_to": body.return_to} + ) + return MoyaiConnectStartResponse(connect_url=connect_url) + + +async def _claim_connect_nonce(prisma_client: "PrismaClient", nonce: str, exp: int) -> None: + from prisma.errors import UniqueViolationError + + try: + await ConfigRepository(prisma_client, use_writer=True).table.create( + data={ + "param_name": f"{_MOYAI_NONCE_CONFIG_PREFIX}{nonce}", + "param_value": json.dumps({"exp": exp}), + } + ) + except UniqueViolationError: + raise HTTPException(status_code=400, detail="Invalid Moyai connect code") + + +async def _moyai_key_alias(prisma_client, moyai_url: str) -> str: + from litellm.repositories.verification_token_repository import ( + VerificationTokenRepository, + ) + + host: Final = urlparse(moyai_url).hostname or "deployment" + alias: Final = f"moyai-{host}" + rows: Final = await VerificationTokenRepository(prisma_client).find_many(where={"key_alias": alias}, take=1) + if rows: + return f"{alias}-{secrets.token_hex(2)}" + return alias + + +async def _persist_moyai_url(prisma_client, moyai_url: str) -> None: + from litellm.proxy.proxy_server import user_api_key_cache + + existing: dict = {} + db_existing: Final = await _ui_settings_db(UISettingsRepository(prisma_client)).find_unique( + where={"id": "ui_settings"} + ) + if db_existing and db_existing.ui_settings: + raw: Final = db_existing.ui_settings + existing = json.loads(raw) if isinstance(raw, str) else dict(raw) + + ui_settings: Final = {**existing, "moyai_url": moyai_url} + await _ui_settings_db(UISettingsRepository(prisma_client)).upsert( + where={"id": "ui_settings"}, + data={ + "create": {"id": "ui_settings", "ui_settings": json.dumps(ui_settings)}, + "update": {"ui_settings": json.dumps(ui_settings)}, + }, + ) + await user_api_key_cache.async_set_cache(key=UI_SETTINGS_CACHE_KEY, value=ui_settings, ttl=UI_SETTINGS_CACHE_TTL) + + +@router.post( + _MOYAI_CONNECT_EXCHANGE_ROUTE, + response_model=MoyaiConnectExchangeResponse, + tags=["moyai"], +) +async def moyai_connect_exchange(request: Request, body: MoyaiConnectExchangeRequest) -> MoyaiConnectExchangeResponse: + from litellm.proxy.management_endpoints.key_management_endpoints import generate_key_helper_fn + from litellm.proxy.proxy_server import llm_router, master_key, prisma_client + + if not master_key: + raise HTTPException(status_code=400, detail="Invalid Moyai connect code") + + payload: Final = _decode_connect_code(master_key, body.code) + + try: + moyai_url: Final = normalize_moyai_url(body.moyai_url) + except ValueError: + raise HTTPException(status_code=400, detail="Invalid Moyai connect code") + if moyai_url is None or _origin(moyai_url) != payload["moyai_origin"]: + raise HTTPException(status_code=400, detail="Invalid Moyai connect code") + + if prisma_client is None: + raise HTTPException(status_code=400, detail="Moyai quick connect needs a database connected to the proxy") + + await _claim_connect_nonce(prisma_client, payload["nonce"], payload["exp"]) + + alias: Final = await _moyai_key_alias(prisma_client, moyai_url) + key_response: Final = await generate_key_helper_fn( + request_type="key", + key_alias=alias, + allowed_routes=_MOYAI_KEY_ALLOWED_ROUTES, + metadata={ + "created_via": "moyai_quick_connect", + "moyai_url": moyai_url, + "connected_by": payload.get("user_id"), + }, + table_name="key", + llm_router=llm_router, + ) + + await _persist_moyai_url(prisma_client, moyai_url) + + return MoyaiConnectExchangeResponse( + api_key=key_response["token"], + key_alias=alias, + api_base=_gateway_url(request), + ) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 8864dc518ab..aa00bfe5d26 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -770,6 +770,7 @@ from litellm.proxy.middleware.request_size_limit_middleware import ( from litellm.proxy.middleware.security_headers_middleware import ( SecurityHeadersMiddleware, ) +from litellm.proxy.moyai_endpoints import router as moyai_router from litellm.proxy.ocr_endpoints.endpoints import router as ocr_router from litellm.proxy.openai_files_endpoints.files_endpoints import ( router as openai_files_router, @@ -20183,6 +20184,7 @@ app.include_router(callback_management_endpoints_router) app.include_router(debugging_endpoints_router) app.include_router(rust_control_plane_router) app.include_router(ui_crud_endpoints_router) +app.include_router(moyai_router) app.include_router(user_banner_endpoints_router) app.include_router(latest_release_endpoints_router) app.include_router(team_callback_router) diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index d307805427a..368c2e9758e 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -15,7 +15,7 @@ from typing import ( from urllib.parse import urlparse from fastapi import APIRouter, Body, Depends, File, HTTPException, UploadFile -from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError, create_model +from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError, create_model, field_validator from pydantic.fields import FieldInfo, PydanticUndefined from typing_extensions import NotRequired, ReadOnly, TypedDict @@ -234,6 +234,20 @@ class UIThemeSettingsResponse(SettingsResponse): _TEAM_ADMIN_FIELD_ENUM: Final = tuple(sorted(SUPPORTED_TEAM_ADMIN_PERMISSIONS)) +def normalize_moyai_url(value: object) -> str | None: + if value is None: + return None + if not isinstance(value, str): + raise ValueError("moyai_url must be a string") + stripped: Final = value.strip() + if not stripped: + return None + parsed: Final = urlparse(stripped) + if parsed.scheme not in ("http", "https") or not parsed.hostname or parsed.username or parsed.password: + raise ValueError("moyai_url must be an http or https URL with a host and no credentials") + return stripped.rstrip("/") + + class UISettings(LiteLLMBaseModel): """Configuration for UI-specific flags""" @@ -330,6 +344,16 @@ class UISettings(LiteLLMBaseModel): description="If true, shows the Chat page in the UI sidebar, letting users chat with an LLM and connect their own MCP server credentials via OAuth.", ) + moyai_url: str | None = Field( + default=None, + description="URL of a connected Moyai deployment. When set, the Moyai entry in the UI navigation opens this deployment instead of the Moyai landing page.", + ) + + @field_validator("moyai_url", mode="before") + @classmethod + def _validate_moyai_url(cls, value: object) -> object: + return normalize_moyai_url(value) + team_admin_editable_team_fields: Sequence[str] = Field( default=(), description=( @@ -368,6 +392,7 @@ ALLOWED_UI_SETTINGS_FIELDS: Final = { "disable_custom_api_keys", "disable_key_generate_for_org_admin", "enable_chat_ui", + "moyai_url", TEAM_ADMIN_EDITABLE_TEAM_FIELDS_SETTING, } diff --git a/tests/unit/proxy/test_moyai_endpoints.py b/tests/unit/proxy/test_moyai_endpoints.py new file mode 100644 index 00000000000..a0b0d56b637 --- /dev/null +++ b/tests/unit/proxy/test_moyai_endpoints.py @@ -0,0 +1,290 @@ +import json +import time +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + + +def _admin() -> UserAPIKeyAuth: + return UserAPIKeyAuth(user_id="admin-user", user_role=LitellmUserRoles.PROXY_ADMIN) + + +def _request() -> MagicMock: + request = MagicMock() + request.base_url = "http://localhost:4000/" + return request + + +@pytest.mark.asyncio +async def test_start_rejects_non_admin(monkeypatch: pytest.MonkeyPatch) -> None: + from fastapi import HTTPException + from litellm.proxy import proxy_server + from litellm.proxy.moyai_endpoints import MoyaiConnectStartRequest, moyai_connect_start + + monkeypatch.setattr(proxy_server, "master_key", "sk-master") + actor = UserAPIKeyAuth(user_id="member", user_role="internal_user") + + with pytest.raises(HTTPException) as exc: + await moyai_connect_start( + _request(), + MoyaiConnectStartRequest(moyai_url="https://moyai.example.com", return_to="http://localhost:3000/ui/moyai"), + actor, + ) + assert exc.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_start_requires_master_key(monkeypatch: pytest.MonkeyPatch) -> None: + from fastapi import HTTPException + from litellm.proxy import proxy_server + from litellm.proxy.moyai_endpoints import MoyaiConnectStartRequest, moyai_connect_start + + monkeypatch.setattr(proxy_server, "master_key", None) + + with pytest.raises(HTTPException) as exc: + await moyai_connect_start( + _request(), + MoyaiConnectStartRequest(moyai_url="https://moyai.example.com", return_to="http://localhost:3000/ui/moyai"), + _admin(), + ) + assert exc.value.status_code == 400 + assert "LITELLM_MASTER_KEY" in exc.value.detail + + +@pytest.mark.asyncio +async def test_start_returns_connect_url_with_all_params(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + from litellm.proxy.moyai_endpoints import MoyaiConnectStartRequest, moyai_connect_start + from urllib.parse import parse_qs, urlparse + + monkeypatch.setattr(proxy_server, "master_key", "sk-master") + + response = await moyai_connect_start( + _request(), + MoyaiConnectStartRequest(moyai_url="https://moyai.example.com/", return_to="http://localhost:3000/ui/moyai"), + _admin(), + ) + + parsed = urlparse(response.connect_url) + assert f"{parsed.scheme}://{parsed.netloc}{parsed.path}" == "https://moyai.example.com/connect/litellm" + params = parse_qs(parsed.query) + assert params["gateway_url"] == ["http://localhost:4000"] + assert params["return_to"] == ["http://localhost:3000/ui/moyai"] + assert params["code"] and "." in params["code"][0] + + +async def _exchange_env(monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + + cache: dict = {} + + async def _get(key: str): + return cache.get(key) + + async def _set(key: str, value, ttl=None): + cache[key] = value + + monkeypatch.setattr(proxy_server, "master_key", "sk-master") + monkeypatch.setattr(proxy_server, "llm_router", None) + user_api_key_cache = SimpleNamespace( + async_get_cache=AsyncMock(side_effect=_get), async_set_cache=AsyncMock(side_effect=_set) + ) + monkeypatch.setattr(proxy_server, "user_api_key_cache", user_api_key_cache) + + prisma = MagicMock() + prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=SimpleNamespace(ui_settings={})) + prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) + + claimed_nonces: set = set() + + async def _config_create(*, data): + from prisma.errors import UniqueViolationError + + param_name = data["param_name"] + if param_name in claimed_nonces: + raise UniqueViolationError({}, message="Unique constraint failed on the fields: (`param_name`)") + claimed_nonces.add(param_name) + + nonce_create = AsyncMock(side_effect=_config_create) + config_table = SimpleNamespace(create=nonce_create) + prisma.db.litellm_config = config_table + prisma.writer_db.litellm_config = config_table + + persisted: dict = {} + + async def _upsert(where, data): + persisted.update(json.loads(data["update"]["ui_settings"])) + + prisma.db.litellm_uisettings.upsert = AsyncMock(side_effect=_upsert) + monkeypatch.setattr(proxy_server, "prisma_client", prisma) + + mint_calls: list = [] + + async def _mint(request_type, **kwargs): + mint_calls.append(kwargs) + return {"token": "sk-new-virtual-key"} + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.generate_key_helper_fn", + AsyncMock(side_effect=_mint), + ) + return persisted, mint_calls, nonce_create + + +async def _exchange_call(code: str, moyai_url: str): + from litellm.proxy.moyai_endpoints import MoyaiConnectExchangeRequest, moyai_connect_exchange + + return await moyai_connect_exchange(_request(), MoyaiConnectExchangeRequest(code=code, moyai_url=moyai_url)) + + +@pytest.mark.asyncio +async def test_exchange_happy_path_mints_key_and_saves_setting(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy.moyai_endpoints import _sign_connect_code + + persisted, mint_calls, _ = await _exchange_env(monkeypatch) + code = _sign_connect_code("sk-master", "https://moyai.example.com", "admin-user") + + response = await _exchange_call(code, "https://moyai.example.com") + + assert response.api_key == "sk-new-virtual-key" + assert response.key_alias == "moyai-moyai.example.com" + assert response.api_base == "http://localhost:4000" + assert persisted["moyai_url"] == "https://moyai.example.com" + mint = mint_calls[0] + assert "user_id" not in mint + assert mint["allowed_routes"] == ["openai_routes", "anthropic_routes", "/model/info"] + assert mint["metadata"] == { + "created_via": "moyai_quick_connect", + "moyai_url": "https://moyai.example.com", + "connected_by": "admin-user", + } + + +@pytest.mark.asyncio +async def test_exchange_without_database_fails_before_nonce(monkeypatch: pytest.MonkeyPatch) -> None: + from fastapi import HTTPException + from litellm.proxy import proxy_server + from litellm.proxy.moyai_endpoints import _sign_connect_code + + _, _, nonce_create = await _exchange_env(monkeypatch) + monkeypatch.setattr(proxy_server, "prisma_client", None) + code = _sign_connect_code("sk-master", "https://moyai.example.com", "admin-user") + + with pytest.raises(HTTPException) as exc: + await _exchange_call(code, "https://moyai.example.com") + assert exc.value.status_code == 400 + assert "database" in exc.value.detail + nonce_create.assert_not_called() + + +@pytest.mark.asyncio +async def test_exchange_rejects_tampered_signature(monkeypatch: pytest.MonkeyPatch) -> None: + from fastapi import HTTPException + from litellm.proxy.moyai_endpoints import _sign_connect_code + + await _exchange_env(monkeypatch) + code = _sign_connect_code("sk-master", "https://moyai.example.com", "admin-user") + tampered = code[:-2] + ("aa" if not code.endswith("aa") else "bb") + + with pytest.raises(HTTPException) as exc: + await _exchange_call(tampered, "https://moyai.example.com") + assert exc.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_exchange_rejects_expired_code(monkeypatch: pytest.MonkeyPatch) -> None: + import base64 + import hashlib + import hmac as hmac_mod + from fastapi import HTTPException + from litellm.proxy.moyai_endpoints import _b64url, _master_key_hmac_key + + await _exchange_env(monkeypatch) + payload = json.dumps( + { + "moyai_origin": "https://moyai.example.com", + "user_id": "admin-user", + "exp": int(time.time()) - 10, + "nonce": "n", + }, + separators=(",", ":"), + sort_keys=True, + ).encode() + sig = hmac_mod.new(_master_key_hmac_key("sk-master"), payload, hashlib.sha256).digest() + code = f"{_b64url(payload)}.{_b64url(sig)}" + + with pytest.raises(HTTPException) as exc: + await _exchange_call(code, "https://moyai.example.com") + assert exc.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_exchange_rejects_origin_mismatch(monkeypatch: pytest.MonkeyPatch) -> None: + from fastapi import HTTPException + from litellm.proxy.moyai_endpoints import _sign_connect_code + + await _exchange_env(monkeypatch) + code = _sign_connect_code("sk-master", "https://moyai.example.com", "admin-user") + + with pytest.raises(HTTPException) as exc: + await _exchange_call(code, "https://evil.example.com") + assert exc.value.status_code == 400 + + +@pytest.mark.asyncio +async def test_exchange_rejects_replayed_nonce(monkeypatch: pytest.MonkeyPatch) -> None: + from fastapi import HTTPException + from litellm.proxy.moyai_endpoints import _sign_connect_code + + _, mint_calls, _ = await _exchange_env(monkeypatch) + code = _sign_connect_code("sk-master", "https://moyai.example.com", "admin-user") + + await _exchange_call(code, "https://moyai.example.com") + with pytest.raises(HTTPException) as exc: + await _exchange_call(code, "https://moyai.example.com") + assert exc.value.status_code == 400 + assert len(mint_calls) == 1 + + +@pytest.mark.asyncio +async def test_exchange_concurrent_replay_claims_nonce_once(monkeypatch: pytest.MonkeyPatch) -> None: + import asyncio + + from fastapi import HTTPException + from litellm.proxy.moyai_endpoints import _sign_connect_code + + _, mint_calls, _ = await _exchange_env(monkeypatch) + code = _sign_connect_code("sk-master", "https://moyai.example.com", "admin-user") + + results = await asyncio.gather( + _exchange_call(code, "https://moyai.example.com"), + _exchange_call(code, "https://moyai.example.com"), + return_exceptions=True, + ) + + successes = [r for r in results if not isinstance(r, BaseException)] + rejections = [r for r in results if isinstance(r, HTTPException) and r.status_code == 400] + assert len(successes) == 1 + assert len(rejections) == 1 + assert len(mint_calls) == 1 + + +@pytest.mark.asyncio +async def test_exchange_replay_survives_fresh_worker_cache(monkeypatch: pytest.MonkeyPatch) -> None: + from fastapi import HTTPException + from litellm.caching.caching import DualCache + from litellm.proxy import proxy_server + from litellm.proxy.moyai_endpoints import _sign_connect_code + + _, mint_calls, _ = await _exchange_env(monkeypatch) + code = _sign_connect_code("sk-master", "https://moyai.example.com", "admin-user") + + await _exchange_call(code, "https://moyai.example.com") + monkeypatch.setattr(proxy_server, "user_api_key_cache", DualCache()) + with pytest.raises(HTTPException) as exc: + await _exchange_call(code, "https://moyai.example.com") + assert exc.value.status_code == 400 + assert len(mint_calls) == 1 diff --git a/tests/unit/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/unit/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index bea305417f3..4a7cee80c61 100644 --- a/tests/unit/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/unit/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -4165,3 +4165,71 @@ class TestSyncUiSettingsToGeneralSettings: assert general_settings["forward_client_headers_to_llm_api"] is False assert general_settings.source("forward_client_headers_to_llm_api") == "config" + + +class TestMoyaiUrlSetting: + @pytest.mark.parametrize( + "value,expected", + [ + (None, None), + ("", None), + ("https://moyai.example.com", "https://moyai.example.com"), + ("https://moyai.example.com/", "https://moyai.example.com"), + ("http://localhost:8787/", "http://localhost:8787"), + ], + ) + def test_moyai_url_validator_accepts_and_normalizes(self, value, expected): + from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import UISettings + + assert UISettings(moyai_url=value).moyai_url == expected + + @pytest.mark.parametrize( + "value", + [ + "javascript:alert(1)", + "ftp://moyai.example.com", + "https://user:pass@moyai.example.com", + "https://user@moyai.example.com", + "not-a-url", + "https://", + ], + ) + def test_moyai_url_validator_rejects(self, value): + from pydantic import ValidationError + from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import UISettings + + with pytest.raises(ValidationError): + UISettings(moyai_url=value) + + def test_moyai_url_is_in_allowed_ui_settings_fields(self): + from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import ALLOWED_UI_SETTINGS_FIELDS + + assert "moyai_url" in ALLOWED_UI_SETTINGS_FIELDS + + @pytest.mark.asyncio + async def test_moyai_url_patch_sets_and_clears(self, monkeypatch): + from types import SimpleNamespace + from unittest.mock import AsyncMock, MagicMock + from litellm.proxy import proxy_server + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.ui_crud_endpoints.proxy_setting_endpoints import update_ui_settings + + prisma = MagicMock() + prisma.db.litellm_uisettings.find_unique = AsyncMock( + return_value=SimpleNamespace(ui_settings={"moyai_url": "https://old.example.com"}) + ) + persisted: dict = {} + + async def _upsert(where, data): + persisted.update(json.loads(data["update"]["ui_settings"])) + + prisma.db.litellm_uisettings.upsert = AsyncMock(side_effect=_upsert) + monkeypatch.setattr(proxy_server, "prisma_client", prisma) + monkeypatch.setattr(proxy_server, "store_model_in_db", True) + actor = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + await update_ui_settings({"moyai_url": "https://new.example.com/"}, actor) + assert persisted["moyai_url"] == "https://new.example.com" + + await update_ui_settings({"moyai_url": None}, actor) + assert persisted["moyai_url"] is None diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/anthropic.svg b/ui/litellm-dashboard/public/assets/moyai/logos/anthropic.svg new file mode 100644 index 00000000000..a37f591fb76 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/anthropic.svg @@ -0,0 +1,5 @@ + + + + + diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/bedrock.svg b/ui/litellm-dashboard/public/assets/moyai/logos/bedrock.svg new file mode 100644 index 00000000000..e0f929a7a97 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/bedrock.svg @@ -0,0 +1 @@ +Bedrock \ No newline at end of file diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/deepseek.svg b/ui/litellm-dashboard/public/assets/moyai/logos/deepseek.svg new file mode 100644 index 00000000000..c4754047da2 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/deepseek.svg @@ -0,0 +1,25 @@ + + + + + + diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/fireworks.svg b/ui/litellm-dashboard/public/assets/moyai/logos/fireworks.svg new file mode 100644 index 00000000000..a23445cf94b --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/fireworks.svg @@ -0,0 +1 @@ +Fireworks \ No newline at end of file diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/google.svg b/ui/litellm-dashboard/public/assets/moyai/logos/google.svg new file mode 100644 index 00000000000..7bc4a38ce7a --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/google.svg @@ -0,0 +1,2 @@ + + \ No newline at end of file diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/hermes.png b/ui/litellm-dashboard/public/assets/moyai/logos/hermes.png new file mode 100644 index 0000000000000000000000000000000000000000..de47b728d12a648a0d3affd06e7d3ed75eacd562 GIT binary patch literal 1100 zcmV-S1he~zP)q7y%f8W&ke&%s_tx@FGChvw#-?7=i6_R^alj?rfw76x79^ zx##o|X}Y_ps(Y&d?m#jmnUXy7&*6OA)5jD@3X%o?dlV$a*x?obKUAF&&dq=Vz_fdM zXxkQT+oG;(e0_bPZCfa%pp=5v8lLCjO_Yw~EW4P$_FGXfj$_7g%-Cx>x!dgqb5cq+P4ki+n40OkR!aGBxYn9gRWVId=6Oz%lu}YkNze0G6a~{XeQCt6 zvFtF$004p@u!-mUK0MDuS(Z>rVY}U8zu)6ola=m7MIHfQcBcyZRe)&F(ot|T-UYNv)K$%%F&HZr_&(*LWqHJQ+P$LB91ix`n-_L=fC~Um7)|O&T(=Z4EDy6vHZYSy5 zr)>Yas;c*L!1sNmX^PEe17i%bEVE3=vJ52Qx~`3HAG6sE^Z9)AMaP~W^*|q5YerG@ z;Gw?OYpo}l(f2^7nD=;UKoA5k3rk0}u_1)8)b%{?y%0>5;3-VH_HLaq#wzSn8uoEK z9$}0@UDu;VJUe`hRbzCVG@!G#TQqi9|6W%&#=OvJNk`wON!V4dj*e5Y>r+rl`Ouhj zvF}T{5aM|bWLaii*1erE2K)W~sh8&SxvhL1uGi}z4;G7sZRI9IV2rUY?oS zC6>!2=JPpLtJPq0DdmF!V^ZBWdNc%46j2Casr8KGNs=%OLjqtJhW7f^NRs412ikrq zUa!~GT64GCQ7JXr4RtZj^ZaRvzSo2i1I(Z)ZJLH!Yx=%#7p&Lohs}Z@u-aZ#)ufSc zcp&=)AOo=OI9rw_4u=D>EJGNE`2G9$Xgrltmgm0j+xt>VD5Y#CWQ-Y9!!Qg7dEezh zIYlw8HG?1+ZONl3Vv;0mnubvnjke-xnp!Vvt@~5u{ij(y#rHm!`@BD&&#bD7d7eKc zS3j(2nsTvN^gHCb`!xHwH688k)9ib!)w-8q&+`z5Aw17RUDv4V`rfqzDej)XY0NhO z|90lTI^Z}?>o|@NULangChain \ No newline at end of file diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/mistral.svg b/ui/litellm-dashboard/public/assets/moyai/logos/mistral.svg new file mode 100644 index 00000000000..8e03e244bf1 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/mistral.svg @@ -0,0 +1 @@ +Mistral \ No newline at end of file diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/openai.svg b/ui/litellm-dashboard/public/assets/moyai/logos/openai.svg new file mode 100644 index 00000000000..52dad8269ec --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/openai.svg @@ -0,0 +1,5 @@ + + + + + \ No newline at end of file diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/opencode.svg b/ui/litellm-dashboard/public/assets/moyai/logos/opencode.svg new file mode 100644 index 00000000000..7ed0af003bb --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/opencode.svg @@ -0,0 +1,16 @@ + + + + + + + + + + + + + + + + diff --git a/ui/litellm-dashboard/public/assets/moyai/logos/xai.svg b/ui/litellm-dashboard/public/assets/moyai/logos/xai.svg new file mode 100644 index 00000000000..9491b192fd5 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/logos/xai.svg @@ -0,0 +1,28 @@ + + + + + + + + + + + + + diff --git a/ui/litellm-dashboard/public/assets/moyai/moyai-head.svg b/ui/litellm-dashboard/public/assets/moyai/moyai-head.svg new file mode 100644 index 00000000000..6d69078c4e3 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/moyai/moyai-head.svg @@ -0,0 +1,19 @@ + + + + + + + + + + + + + + + + + + + diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableBlogPosts.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableBlogPosts.ts index a7b37b78d42..8bffbe9f741 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableBlogPosts.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableBlogPosts.ts @@ -29,5 +29,5 @@ function getSnapshot() { } export function useDisableBlogPosts() { - return useSyncExternalStore(subscribe, getSnapshot); + return useSyncExternalStore(subscribe, getSnapshot, () => false); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableBouncingIcon.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableBouncingIcon.ts index f5d8087ebe7..0ec94f9915f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableBouncingIcon.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableBouncingIcon.ts @@ -29,5 +29,5 @@ function getSnapshot() { } export function useDisableBouncingIcon() { - return useSyncExternalStore(subscribe, getSnapshot); + return useSyncExternalStore(subscribe, getSnapshot, () => false); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowNewBadge.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowNewBadge.ts index d0a618e27ba..0feba3a27fd 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowNewBadge.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowNewBadge.ts @@ -31,5 +31,5 @@ function getSnapshot() { } export function useDisableShowNewBadge() { - return useSyncExternalStore(subscribe, getSnapshot); + return useSyncExternalStore(subscribe, getSnapshot, () => false); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowPrompts.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowPrompts.ts index 801fbdbb99d..872da097868 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowPrompts.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useDisableShowPrompts.ts @@ -31,5 +31,5 @@ function getSnapshot() { } export function useDisableShowPrompts() { - return useSyncExternalStore(subscribe, getSnapshot); + return useSyncExternalStore(subscribe, getSnapshot, () => false); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useHideAutoRouterAnnouncement.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useHideAutoRouterAnnouncement.ts index c213c4834bc..61b9b3eea7a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/useHideAutoRouterAnnouncement.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/useHideAutoRouterAnnouncement.ts @@ -31,5 +31,5 @@ function getSnapshot() { } export function useHideAutoRouterAnnouncement() { - return useSyncExternalStore(subscribe, getSnapshot); + return useSyncExternalStore(subscribe, getSnapshot, () => false); } diff --git a/ui/litellm-dashboard/src/app/moyai/page.tsx b/ui/litellm-dashboard/src/app/moyai/page.tsx new file mode 100644 index 00000000000..0c2179714b5 --- /dev/null +++ b/ui/litellm-dashboard/src/app/moyai/page.tsx @@ -0,0 +1,99 @@ +"use client"; + +import { Suspense, useEffect, useMemo, useSyncExternalStore } from "react"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings"; +import Navbar from "@/components/navbar"; +import MoyaiConnected from "@/components/moyai/MoyaiConnected"; +import MoyaiLanding from "@/components/moyai/MoyaiLanding"; +import { startMoyaiQuickConnect } from "@/components/networking"; +import { PluginModeProvider } from "@/contexts/PluginModeContext"; +import { ThemeProvider } from "@/contexts/ThemeContext"; +import { isProxyAdminRole } from "@/utils/roles"; +import { uiHref } from "@/utils/uiHref"; + +function MoyaiPageContent() { + const { accessToken, userRole } = useAuthorized(); + const { data: uiSettings, isLoading, refetch } = useUISettings(); + + const parsed = useSyncExternalStore( + () => () => {}, + () => true, + () => false, + ); + const connectedParams = useMemo(() => { + if (!parsed) { + return null; + } + const params = new URLSearchParams(window.location.search); + if (params.get("moyai_connected") !== "1") { + return null; + } + const models = Number(params.get("models")); + return { + keyAlias: params.get("key_alias"), + models: Number.isFinite(models) && params.get("models") !== null ? models : null, + }; + }, [parsed]); + + useEffect(() => { + if (connectedParams) { + window.history.replaceState(null, "", window.location.pathname); + refetch(); + } + }, [connectedParams, refetch]); + + const moyaiUrl = (uiSettings?.values?.moyai_url as string | undefined) ?? null; + + const settingsReady = parsed && !isLoading; + const shouldOpenMoyai = settingsReady && !connectedParams && moyaiUrl; + useEffect(() => { + if (shouldOpenMoyai) { + window.location.replace(moyaiUrl as string); + } + }, [shouldOpenMoyai, moyaiUrl]); + + const onQuickConnect = async (url: string) => { + const response = await startMoyaiQuickConnect(accessToken ?? "", url, window.location.origin + uiHref("moyai")); + window.location.assign(response.connect_url); + }; + + let content: React.ReactNode = null; + if (parsed && !isLoading) { + if (connectedParams) { + content = ( + + ); + } else if (moyaiUrl) { + content = ( +
+

Opening Moyai

+ + Continue to {moyaiUrl} + +
+ ); + } else { + content = ; + } + } + + return ( + + +
+ +
{content}
+
+
+
+ ); +} + +export default function MoyaiPage() { + return ( + + + + ); +} diff --git a/ui/litellm-dashboard/src/components/Navbar/ViewSwitcher.test.tsx b/ui/litellm-dashboard/src/components/Navbar/ViewSwitcher.test.tsx index a32df932d05..882623b7de7 100644 --- a/ui/litellm-dashboard/src/components/Navbar/ViewSwitcher.test.tsx +++ b/ui/litellm-dashboard/src/components/Navbar/ViewSwitcher.test.tsx @@ -9,6 +9,7 @@ const { mockUsePluginMode, mockUseUISettings, mockUsePathname, state } = vi.hois plugins: [] as { name: string; display_name: string; url: string }[], activePlugin: null as { name: string; display_name: string; url: string } | null, enableChatUI: false, + moyaiUrl: undefined as string | undefined, pathname: "/ui/", }; return { @@ -19,7 +20,9 @@ const { mockUsePluginMode, mockUseUISettings, mockUsePathname, state } = vi.hois plugins: state.plugins, activePlugin: state.activePlugin, })), - mockUseUISettings: vi.fn(() => ({ data: { values: { enable_chat_ui: state.enableChatUI } } })), + mockUseUISettings: vi.fn(() => ({ + data: { values: { enable_chat_ui: state.enableChatUI, moyai_url: state.moyaiUrl } }, + })), mockUsePathname: vi.fn(() => state.pathname), }; }); @@ -45,6 +48,7 @@ describe("ViewSwitcher", () => { state.mode = "ai-gateway"; state.plugins = []; state.enableChatUI = false; + state.moyaiUrl = undefined; state.pathname = "/ui/"; state.setMode.mockClear(); }); @@ -135,6 +139,81 @@ describe("ViewSwitcher", () => { expect(assignSpy).toHaveBeenCalledWith("/ui/"); }); + it("shows the Moyai entry with its description regardless of enable_chat_ui", async () => { + render(); + + act(() => { + fireEvent.click(screen.getByRole("button")); + }); + expect(await screen.findByText("Moyai")).toBeInTheDocument(); + expect(screen.getByText("Cloud Coding Agent")).toBeInTheDocument(); + }); + + it("navigates to the moyai route when the Moyai entry is picked", async () => { + render(); + + act(() => { + fireEvent.click(screen.getByRole("button")); + }); + act(() => { + fireEvent.click(screen.getByText("Moyai")); + }); + expect(assignSpy).toHaveBeenCalledWith("/ui/moyai"); + expect(state.setMode).not.toHaveBeenCalled(); + }); + + it("labels the button Moyai on the moyai route and its subpaths", async () => { + state.pathname = "/ui/moyai"; + const { unmount } = render(); + expect(screen.getByRole("button")).toHaveTextContent("Moyai"); + unmount(); + + state.pathname = "/ui/moyai/anything"; + render(); + expect(screen.getByRole("button")).toHaveTextContent("Moyai"); + }); + + it("navigates back to the dashboard when AI Gateway is picked from the moyai route", async () => { + state.pathname = "/ui/moyai"; + render(); + + act(() => { + fireEvent.click(screen.getByRole("button")); + }); + expect(await screen.findByText("AI Gateway")).toBeInTheDocument(); + act(() => { + fireEvent.click(screen.getByText("AI Gateway")); + }); + expect(state.setMode).toHaveBeenCalledWith("ai-gateway"); + expect(assignSpy).toHaveBeenCalledWith("/ui/"); + }); + + it("lists Moyai before Chat in the menu", async () => { + state.enableChatUI = true; + render(); + + act(() => { + fireEvent.click(screen.getByRole("button")); + }); + expect(await screen.findByText("Moyai")).toBeInTheDocument(); + expect(screen.getByText("Moyai").compareDocumentPosition(screen.getByText("Chat"))).toBe( + Node.DOCUMENT_POSITION_FOLLOWING, + ); + }); + + it("navigates straight to the connected deployment when moyai_url is set", async () => { + state.moyaiUrl = "https://moyai.example.com"; + render(); + + act(() => { + fireEvent.click(screen.getByRole("button")); + }); + act(() => { + fireEvent.click(screen.getByText("Moyai")); + }); + expect(assignSpy).toHaveBeenCalledWith("https://moyai.example.com"); + }); + it("shows Chat as a disabled, non-navigating entry with an admin hint when disabled", async () => { state.enableChatUI = false; state.plugins = [{ name: "obs", display_name: "Observability", url: "http://localhost:9000" }]; diff --git a/ui/litellm-dashboard/src/components/Navbar/ViewSwitcher.tsx b/ui/litellm-dashboard/src/components/Navbar/ViewSwitcher.tsx index b3aba7155d1..24d9e332ad4 100644 --- a/ui/litellm-dashboard/src/components/Navbar/ViewSwitcher.tsx +++ b/ui/litellm-dashboard/src/components/Navbar/ViewSwitcher.tsx @@ -10,9 +10,21 @@ import { Check, ChevronsUpDown, LayoutGrid } from "lucide-react"; import { usePluginMode } from "@/contexts/PluginModeContext"; import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings"; import { uiHref } from "@/utils/uiHref"; +import moyaiHead from "../../../public/assets/moyai/moyai-head.svg"; const GATEWAY = "ai-gateway"; const CHAT = "chat"; +const MOYAI = "moyai"; + +function isRoute(pathname: string, href: string): boolean { + return pathname === href || pathname.startsWith(`${href}/`); +} + +function activeLabelFor(isChatRoute: boolean, isMoyaiRoute: boolean, pluginLabel: string | undefined): string { + if (isChatRoute) return "Chat"; + if (isMoyaiRoute) return "Moyai"; + return pluginLabel ?? "AI Gateway"; +} interface ViewSwitcherItem { key: string; @@ -27,12 +39,14 @@ export default function ViewSwitcher() { const pathname = usePathname(); const chatEnabled = Boolean(uiSettings?.values?.enable_chat_ui); + const moyaiUrl = (uiSettings?.values?.moyai_url as string | undefined) ?? null; - const chatHref = uiHref(CHAT); const normalizedPathname = (pathname ?? "").replace(/\/+$/, ""); - const isChatRoute = chatEnabled && (normalizedPathname === chatHref || normalizedPathname.startsWith(`${chatHref}/`)); + const isChatRoute = chatEnabled && isRoute(normalizedPathname, uiHref(CHAT)); + const isMoyaiRoute = isRoute(normalizedPathname, uiHref(MOYAI)); + const isStandaloneRoute = isChatRoute || isMoyaiRoute; - const activeLabel = isChatRoute ? "Chat" : plugins.find((p) => p.name === mode)?.display_name ?? "AI Gateway"; + const activeLabel = activeLabelFor(isChatRoute, isMoyaiRoute, plugins.find((p) => p.name === mode)?.display_name); const modeEntries = [ { key: GATEWAY, label: "AI Gateway" }, @@ -41,9 +55,7 @@ export default function ViewSwitcher() { const selectMode = (key: string) => { setMode(key); - // The chat route lives outside the dashboard SPA shell that reacts to `mode`, - // so switching modes from there needs a real navigation, not just state. - if (isChatRoute) { + if (isStandaloneRoute) { window.location.assign(uiHref("")); } }; @@ -78,11 +90,27 @@ export default function ViewSwitcher() { label: (
{e.label} - {!isChatRoute && e.key === mode && } + {!isStandaloneRoute && e.key === mode && }
), onClick: () => selectMode(e.key), })), + { + key: MOYAI, + label: ( +
+ + + + Moyai + Cloud Coding Agent + + + {isMoyaiRoute && } +
+ ), + onClick: () => window.location.assign(moyaiUrl ?? uiHref(MOYAI)), + }, chatItem, ]; diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/UISettings/UISettings.test.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/UISettings/UISettings.test.tsx index c2834e65498..6a491f6160f 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/UISettings/UISettings.test.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/UISettings/UISettings.test.tsx @@ -127,6 +127,45 @@ describe("UISettings", () => { expect(toast.success).toHaveBeenCalledWith("UI settings updated successfully"); }); + it("disconnects Moyai by patching moyai_url to null", () => { + const mutateMock = vi.fn((_settings, options) => { + options?.onSuccess?.(); + }); + + mockUseUpdateUISettings.mockReturnValue({ + mutate: mutateMock, + isPending: false, + error: null, + }); + mockUseUISettings.mockReturnValue( + buildSettingsResponse({ + data: { + ...buildSettingsResponse().data, + values: { ...buildSettingsResponse().data.values, moyai_url: "https://moyai.example.com" }, + }, + }), + ); + + render(); + + expect(screen.getByText("Connected to https://moyai.example.com")).toBeInTheDocument(); + act(() => { + fireEvent.click(screen.getByRole("button", { name: "Disconnect" })); + }); + + expect(mutateMock).toHaveBeenCalledWith( + { moyai_url: null }, + expect.objectContaining({ onSuccess: expect.any(Function) }), + ); + }); + + it("shows a link to the Moyai page when not connected", () => { + render(); + + expect(screen.getByText("Not connected")).toBeInTheDocument(); + expect(screen.getByRole("link", { name: "Connect from the Moyai page" })).toHaveAttribute("href", "/ui/moyai"); + }); + it("should toggle require auth for public AI Hub setting and call update", () => { const mutateMock = vi.fn((_settings, options) => { options?.onSuccess?.(); diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/UISettings/UISettings.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/UISettings/UISettings.tsx index 612ca05d083..476a4f2f6b8 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/UISettings/UISettings.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/UISettings/UISettings.tsx @@ -9,6 +9,8 @@ import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card"; import { Separator } from "@/components/ui/separator"; import { Skeleton } from "@/components/ui/skeleton"; import { Switch } from "@/components/ui/switch"; +import { Button } from "@/components/ui/button"; +import { uiHref } from "@/utils/uiHref"; import PageVisibilitySettings from "./PageVisibilitySettings"; interface SettingRowProps { @@ -66,6 +68,7 @@ export default function UISettings() { const scopeUserSearchProperty = schema?.properties?.scope_user_search_to_org; const disableCustomApiKeysProperty = schema?.properties?.disable_custom_api_keys; const values = data?.values ?? {}; + const moyaiUrl = (values.moyai_url as string | undefined) ?? null; const isDisabledForInternalUsers = Boolean(values.disable_model_add_for_internal_users); const isDisabledTeamAdminDeleteTeamUser = Boolean(values.disable_team_admin_delete_team_user); const isAgentsDisabled = Boolean(values.disable_agents_for_internal_users); @@ -252,6 +255,20 @@ export default function UISettings() { ); }; + const handleDisconnectMoyai = () => { + updateSettings( + { moyai_url: null }, + { + onSuccess: () => { + toast.success("UI settings updated successfully"); + }, + onError: (error) => { + toast.fromError(error); + }, + }, + ); + }; + const handleToggleDisableCustomApiKeys = (checked: boolean) => { updateSettings( { disable_custom_api_keys: checked }, @@ -366,6 +383,25 @@ export default function UISettings() { } /> + +
+
+

Moyai

+

+ {moyaiUrl ? `Connected to ${moyaiUrl}` : "Not connected"} +

+
+ {moyaiUrl ? ( + + ) : ( + + Connect from the Moyai page + + )} +
+ ({ uiHref: (seg: string) => `/ui/${seg}` })); + +describe("MoyaiConnected", () => { + it("shows workspace, key alias, and model count rows with an Open Moyai CTA", () => { + render(); + + expect(screen.getByText("Workspace linked")).toBeInTheDocument(); + expect(screen.getByText("https://moyai.example.com")).toBeInTheDocument(); + expect(screen.getByText("Virtual key issued")).toBeInTheDocument(); + expect(screen.getByText("moyai-moyai.example.com")).toBeInTheDocument(); + expect(screen.getByText(/models available through LiteLLM/)).toHaveTextContent( + "42 models available through LiteLLM", + ); + expect(screen.getByRole("link", { name: /Open Moyai/ })).toHaveAttribute("href", "https://moyai.example.com"); + expect(screen.getByRole("link", { name: /Back to AI Gateway/ })).toHaveAttribute("href", "/ui/"); + }); + + it("hides the models row when the count is missing", () => { + render(); + + expect(screen.queryByText(/models available through LiteLLM/)).not.toBeInTheDocument(); + expect(screen.queryByText("Virtual key issued")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/moyai/MoyaiConnected.tsx b/ui/litellm-dashboard/src/components/moyai/MoyaiConnected.tsx new file mode 100644 index 00000000000..9030e2ab96d --- /dev/null +++ b/ui/litellm-dashboard/src/components/moyai/MoyaiConnected.tsx @@ -0,0 +1,84 @@ +"use client"; + +import React from "react"; +import { ArrowUpRight, Check, ExternalLink } from "lucide-react"; +import moyaiHead from "../../../public/assets/moyai/moyai-head.svg"; +import styles from "./MoyaiLanding.module.css"; +import { uiHref } from "@/utils/uiHref"; + +export default function MoyaiConnected({ + moyaiUrl, + keyAlias, + models, +}: { + moyaiUrl: string | null; + keyAlias?: string | null; + models?: number | null; +}) { + const rows = [ + { label: "Workspace linked", value: moyaiUrl }, + { label: "Virtual key issued", value: keyAlias }, + { label: "models available through LiteLLM", value: models != null ? `${models}` : null, prefix: true }, + ]; + + return ( +
+
+
+ + Connected +
+ +

Moyai connected

+

+ This deployment is now linked to your LiteLLM gateway. +

+ +
+ {rows.map((row) => + row.value ? ( +
+ + + {row.prefix ? ( + <> + {row.value} {row.label} + + ) : ( + <> + {row.label} {row.value} + + )} + +
+ ) : null, + )} +
+ + +
+
+ ); +} diff --git a/ui/litellm-dashboard/src/components/moyai/MoyaiLanding.module.css b/ui/litellm-dashboard/src/components/moyai/MoyaiLanding.module.css new file mode 100644 index 00000000000..cd57b893017 --- /dev/null +++ b/ui/litellm-dashboard/src/components/moyai/MoyaiLanding.module.css @@ -0,0 +1,84 @@ +@keyframes moyai-spin { + to { + --moyai-angle: 360deg; + } +} + +@keyframes moyai-rise { + from { + opacity: 0; + transform: translateY(16px); + } + to { + opacity: 1; + transform: none; + } +} + +@keyframes moyai-fade { + from { + opacity: 0; + } +} + +@property --moyai-angle { + syntax: ""; + initial-value: 0deg; + inherits: false; +} + +.rise { + animation: moyai-rise 0.9s cubic-bezier(0.2, 0.7, 0.2, 1) both; +} + +.orbit { + visibility: hidden; +} + +.orbit[data-placed="true"] { + visibility: visible; + animation: moyai-fade 1.2s ease-out both; +} + +.cta { + background: + linear-gradient(#ffffff, #ffffff) padding-box, + conic-gradient(from var(--moyai-angle), #8b9bff, #cfe3ff, #ffb98a, #e9d8ff, #8b9bff) border-box; + border: 2px solid transparent; + box-shadow: + 0 0 32px rgba(139, 155, 255, 0.55), + 0 0 90px rgba(120, 170, 255, 0.35); + animation: moyai-spin 4s linear infinite; +} + +.cta:hover { + box-shadow: + 0 0 44px rgba(139, 155, 255, 0.8), + 0 0 120px rgba(120, 170, 255, 0.5); +} + +.connect { + background: + linear-gradient(#0b1020, #0b1020) padding-box, + conic-gradient(from var(--moyai-angle), #8b9bff, #cfe3ff, #ffb98a, #e9d8ff, #8b9bff) border-box; + border: 1.5px solid transparent; + box-shadow: 0 0 22px rgba(139, 155, 255, 0.45); + animation: moyai-spin 4s linear infinite; +} + +.connect:hover { + box-shadow: 0 0 34px rgba(139, 155, 255, 0.75); +} + +.mark { + filter: drop-shadow(0 0 0.18em rgba(140, 180, 255, 0.45)) drop-shadow(0 0.04em 0.08em rgba(0, 0, 0, 0.6)); +} + +@media (prefers-reduced-motion: reduce) { + .rise, + .cta, + .connect, + .orbit[data-placed="true"] { + animation: none; + } +} diff --git a/ui/litellm-dashboard/src/components/moyai/MoyaiLanding.test.tsx b/ui/litellm-dashboard/src/components/moyai/MoyaiLanding.test.tsx new file mode 100644 index 00000000000..3d3095ddd5e --- /dev/null +++ b/ui/litellm-dashboard/src/components/moyai/MoyaiLanding.test.tsx @@ -0,0 +1,56 @@ +import { describe, expect, it, vi } from "vitest"; +import { fireEvent, render, screen } from "@testing-library/react"; +import MoyaiLanding, { MOYAI_GITHUB_URL, MOYAI_LAUNCH_POST_URL, MOYAI_WALKTHROUGH_URL } from "./MoyaiLanding"; + +vi.mock("./moyaiSky", async (importOriginal) => { + const actual = await importOriginal(); + return { + ...actual, + prefersReducedMotion: () => true, + startStarfield: () => () => {}, + startPlanetrise: () => () => {}, + }; +}); + +describe("MoyaiLanding", () => { + it("shows the admin quick-connect line and validates the dialog input", async () => { + const onQuickConnect = vi.fn(async () => {}); + render(); + + fireEvent.click(screen.getByRole("button", { name: "Quick connect" })); + + const input = await screen.findByLabelText("Moyai URL"); + fireEvent.change(input, { target: { value: "javascript:alert(1)" } }); + fireEvent.click(screen.getByRole("button", { name: "Connect" })); + expect(await screen.findByRole("alert")).toHaveTextContent("http or https"); + expect(onQuickConnect).not.toHaveBeenCalled(); + + fireEvent.change(input, { target: { value: " https://moyai.example.com/ " } }); + fireEvent.click(screen.getByRole("button", { name: "Connect" })); + expect(onQuickConnect).toHaveBeenCalledWith("https://moyai.example.com"); + }); + + it("shows the non-admin copy when quick connect is unavailable", () => { + render(); + + expect(screen.getByText(/Ask a proxy admin to connect it/)).toBeInTheDocument(); + expect(screen.queryByRole("button", { name: "Quick connect" })).not.toBeInTheDocument(); + }); + + it("links the GitHub, demo, and launch post CTAs to the exported URLs", () => { + render(); + + expect(screen.getByRole("link", { name: /Star Moyai on GitHub/ })).toHaveAttribute("href", MOYAI_GITHUB_URL); + expect(screen.getByRole("link", { name: /Watch the demo/ })).toHaveAttribute("href", MOYAI_WALKTHROUGH_URL); + expect(screen.getByRole("link", { name: /Read the launch post/ })).toHaveAttribute("href", MOYAI_LAUNCH_POST_URL); + }); + + it("shows the GitHub fallback when the demo image fails to load", () => { + render(); + + const demoImg = screen.getByAltText(/Moyai demo:/); + fireEvent.error(demoImg); + + expect(screen.getByText("Watch the demo on GitHub")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/moyai/MoyaiLanding.tsx b/ui/litellm-dashboard/src/components/moyai/MoyaiLanding.tsx new file mode 100644 index 00000000000..e64c153145b --- /dev/null +++ b/ui/litellm-dashboard/src/components/moyai/MoyaiLanding.tsx @@ -0,0 +1,463 @@ +"use client"; + +import React, { useEffect, useRef, useState } from "react"; +import { ArrowRight, ArrowUpRight, Github, Play, Plug, Star } from "lucide-react"; +import moyaiHead from "../../../public/assets/moyai/moyai-head.svg"; +import anthropicLogo from "../../../public/assets/moyai/logos/anthropic.svg"; +import bedrockLogo from "../../../public/assets/moyai/logos/bedrock.svg"; +import deepseekLogo from "../../../public/assets/moyai/logos/deepseek.svg"; +import fireworksLogo from "../../../public/assets/moyai/logos/fireworks.svg"; +import googleLogo from "../../../public/assets/moyai/logos/google.svg"; +import hermesLogo from "../../../public/assets/moyai/logos/hermes.png"; +import langchainLogo from "../../../public/assets/moyai/logos/langchain.svg"; +import mistralLogo from "../../../public/assets/moyai/logos/mistral.svg"; +import openaiLogo from "../../../public/assets/moyai/logos/openai.svg"; +import opencodeLogo from "../../../public/assets/moyai/logos/opencode.svg"; +import xaiLogo from "../../../public/assets/moyai/logos/xai.svg"; +import { PLANET_HORIZON, prefersReducedMotion, startPlanetrise, startStarfield } from "./moyaiSky"; +import { normalizeMoyaiUrl } from "./moyaiConnect"; +import { Dialog, DialogContent, DialogHeader, DialogTitle, DialogTrigger } from "@/components/ui/dialog"; +import styles from "./MoyaiLanding.module.css"; + +export const MOYAI_GITHUB_URL = "https://github.com/BerriAI/moyai"; +export const MOYAI_WALKTHROUGH_URL = `${MOYAI_GITHUB_URL}#see-it-in-action`; +export const MOYAI_LAUNCH_POST_URL = "https://docs.litellm.ai/blog/moyai-open-source"; +const MOYAI_DEMO_GIF_URL = "https://github.com/user-attachments/assets/2de74e6a-c37c-48d6-8a99-de2166c88626"; + +interface LogoItem { + name: string; + logo: { src: string }; +} + +const HARNESSES: LogoItem[] = [ + { name: "Claude Code", logo: anthropicLogo }, + { name: "Codex", logo: openaiLogo }, + { name: "Hermes", logo: hermesLogo }, + { name: "OpenCode", logo: opencodeLogo }, + { name: "Deep Agents", logo: langchainLogo }, +]; + +const PROVIDERS: LogoItem[] = [ + { name: "OpenAI", logo: openaiLogo }, + { name: "Anthropic", logo: anthropicLogo }, + { name: "Fireworks", logo: fireworksLogo }, + { name: "Google", logo: googleLogo }, + { name: "xAI", logo: xaiLogo }, + { name: "Mistral", logo: mistralLogo }, + { name: "DeepSeek", logo: deepseekLogo }, + { name: "Bedrock", logo: bedrockLogo }, +]; + +const ORBIT_TILT = 0.2; +const INNER_RADIUS: [number, number] = [0.16, 240]; +const OUTER_RADIUS: [number, number] = [0.28, 420]; + +function useCanvasScene(start: (canvas: HTMLCanvasElement) => () => void) { + const ref = useRef(null); + useEffect(() => { + if (!ref.current) return; + return start(ref.current); + }, [start]); + return ref; +} + +function Orbit({ + items, + radius, + speed, + heroRef, + anchorRef, + delay, + className = "", +}: { + items: LogoItem[]; + radius: [number, number]; + speed: number; + heroRef: React.RefObject; + anchorRef: React.RefObject; + delay?: string; + className?: string; +}) { + const ref = useRef(null); + useEffect(() => { + const el = ref.current; + if (!el) return; + const reduce = prefersReducedMotion(); + let rafId = 0; + const frame = (ms: number) => { + const hero = heroRef.current?.getBoundingClientRect(); + const anchor = anchorRef.current?.getBoundingClientRect(); + if (hero && anchor) { + const t = ms / 1000; + const top = anchor.bottom - hero.top + 24; + const bottom = hero.height * PLANET_HORIZON - 24; + const outer = Math.min(hero.width * OUTER_RADIUS[0], OUTER_RADIUS[1]); + const scale = Math.min(1, Math.max(0, (bottom - top) / 2 - 18) / (outer * ORBIT_TILT)); + const rx = Math.min(hero.width * radius[0], radius[1]) * scale; + const ry = rx * ORBIT_TILT; + el.style.top = `${(top + bottom) / 2}px`; + el.dataset.placed = "true"; + Array.from(el.children).forEach((child, i) => { + const chip = child as HTMLElement; + const theta = (i / items.length) * Math.PI * 2 + t * speed; + const front = Math.sin(theta) > 0; + chip.style.transform = `translate(${Math.cos(theta) * rx}px, ${Math.sin(theta) * ry}px) translate(-50%, -50%)`; + chip.style.zIndex = front ? "2" : "1"; + chip.style.opacity = front ? "1" : "0.55"; + }); + } + if (!reduce) rafId = requestAnimationFrame(frame); + }; + rafId = requestAnimationFrame(frame); + let resizeObserver: ResizeObserver | null = null; + let onResize: (() => void) | null = null; + if (reduce) { + const reposition = () => frame(performance.now()); + if (typeof ResizeObserver !== "undefined" && heroRef.current) { + resizeObserver = new ResizeObserver(reposition); + resizeObserver.observe(heroRef.current); + } else { + onResize = reposition; + window.addEventListener("resize", reposition); + } + } + return () => { + cancelAnimationFrame(rafId); + resizeObserver?.disconnect(); + if (onResize) window.removeEventListener("resize", onResize); + }; + }, [items, radius, speed, heroRef, anchorRef]); + + return ( + + ); +} + +function LogoWall({ title, items, footnote }: { title: string; items: LogoItem[]; footnote: string }) { + return ( +
+
{title}
+
+ {items.map((item) => ( +
+ + {item.name} +
+ ))} +
+

{footnote}

+
+ ); +} + +function GithubCta({ label = "Star Moyai on GitHub" }: { label?: string }) { + return ( + + + {label} + + + ); +} + +function SecondaryCta({ href, icon, children }: { href: string; icon: React.ReactNode; children: React.ReactNode }) { + return ( + + {icon} + {children} + + ); +} + +function DemoPreview() { + const [failed, setFailed] = useState(false); + return ( + +
+ + + + github.com/BerriAI/moyai +
+ {failed ? ( +
+ + Watch the demo on GitHub +
+ ) : ( + Moyai demo: picking a harness and model, then running a task setFailed(true)} + /> + )} + + Full walkthrough on GitHub + +
+ ); +} + +function QuickConnectDialog({ onQuickConnect }: { onQuickConnect: (url: string) => Promise | void }) { + const [open, setOpen] = useState(false); + const [url, setUrl] = useState(""); + const [error, setError] = useState(null); + const [pending, setPending] = useState(false); + + const connect = async () => { + const normalized = normalizeMoyaiUrl(url); + if (!normalized) { + setError("Enter a full http or https URL, without credentials"); + return; + } + setError(null); + setPending(true); + try { + await onQuickConnect(normalized); + } catch (e) { + setError(e instanceof Error ? e.message : "Could not start quick connect"); + setPending(false); + } + }; + + return ( + + + } + > + + Quick connect + + + + + Quick connect Moyai + +
+ + setUrl(e.target.value)} + placeholder="https://moyai.your-company.com" + className="w-full rounded-md border border-border bg-background px-3 py-2 text-sm" + /> +

+ You’ll confirm on your Moyai workspace as an admin, then come straight back +

+ {error && ( +

+ {error} +

+ )} +
+
+ +
+
+
+ ); +} + +const STATS = [ + { value: "79%", label: "cheaper than our Devin bill" }, + { value: "100+", label: "providers through LiteLLM" }, + { value: "6", label: "agent harnesses" }, +]; + +export default function MoyaiLanding({ + canQuickConnect = false, + onQuickConnect, +}: { + canQuickConnect?: boolean; + onQuickConnect?: (url: string) => Promise | void; +}) { + const heroRef = useRef(null); + const anchorRef = useRef(null); + const starsRef = useCanvasScene(startStarfield); + const planetRef = useCanvasScene(startPlanetrise); + + return ( +
+
- - - - - formatMetric(Number(value), shown)} - /> - - `${label} · Gateway total ${formatMetric(bucketTotals.get(String(label)) ?? 0, shown)}` - } - /> - } - /> - {models.map((model, index) => ( - - ))} - - + formatMetric(value, shown)} + totalLabel="Gateway total" + totalFor={(label) => bucketTotals.get(label) ?? 0} + /> diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/EndpointUsage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/EndpointUsage.tsx index 51e451ca770..dcf248cee65 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/EndpointUsage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/EndpointUsage.tsx @@ -56,10 +56,12 @@ const EndpointUsage: React.FC = ({ userSpendData }) => { }, [userSpendData]); return ( -
+
+
+ + +
- -
); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.test.tsx index a9e65b21f4b..0c9407d6ea1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.test.tsx @@ -30,18 +30,20 @@ describe("EndpointUsageBarChart", () => { renderWithProviders(); expect(screen.getByText("Success vs Failed Requests by Endpoint")).toBeInTheDocument(); - expect(screen.getByText("Successful Requests")).toBeInTheDocument(); - expect(screen.getByText("Failed Requests")).toBeInTheDocument(); + expect(screen.getByText("Successful")).toBeInTheDocument(); + expect(screen.getByText("Failed")).toBeInTheDocument(); + expect(screen.getByText("160")).toBeInTheDocument(); + expect(screen.getByText("7")).toBeInTheDocument(); }); - it("renders stacked green and red bars per endpoint", () => { + it("renders stacked brand-blue and red bars per endpoint", () => { const { container } = renderWithProviders(); expect(container.querySelectorAll(".recharts-bar")).toHaveLength(2); const rectangles = Array.from(container.querySelectorAll("path.recharts-rectangle")); expect(rectangles).toHaveLength(4); const fills = new Set(rectangles.map((rect) => rect.getAttribute("fill"))); - expect(fills).toEqual(new Set(["var(--color-green-500, #22c55e)", "var(--color-red-500, #ef4444)"])); + expect(fills).toEqual(new Set(["#2b3fd6", "#ef4444"])); const xPositions = rectangles.map((rect) => rect.getAttribute("d")?.split(",")[0]); expect(new Set(xPositions).size).toBe(2); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.tsx index bf9868d77cf..b55ea0513c7 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageBarChart.tsx @@ -1,53 +1,69 @@ import React from "react"; -import { BarChart, CustomLegend, CustomTooltip } from "@/components/shared/charts"; -import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card"; +import { ChartColumnStacked } from "lucide-react"; +import { BarChart, CustomTooltip } from "@/components/shared/charts"; import { MetricWithMetadata } from "@/components/UsagePage/types"; +import { Panel } from "../../overview/Primitives"; interface EndpointUsageBarChartProps { endpointData?: Record; } +const SUCCESS_COLOR = "#2b3fd6"; +const FAILED_COLOR = "#ef4444"; +const CATEGORIES = ["Successful", "Failed"] as const; +const COLORS = [SUCCESS_COLOR, FAILED_COLOR] as const; + +const valueFormatter = (value: number) => value.toLocaleString(); + const EndpointUsageBarChart: React.FC = ({ endpointData }) => { // Transform endpoint data into chart format const chartData = React.useMemo(() => { return Object.entries(endpointData || {}).map(([endpoint, data]) => ({ endpoint, - "metrics.successful_requests": data.metrics.successful_requests, - "metrics.failed_requests": data.metrics.failed_requests, - metrics: { - successful_requests: data.metrics.successful_requests, - failed_requests: data.metrics.failed_requests, - }, + Successful: data.metrics.successful_requests, + Failed: data.metrics.failed_requests, })); }, [endpointData]); - const valueFormatter = (value: number) => value.toLocaleString(); + const totals = React.useMemo( + () => + chartData.reduce( + (acc, row) => ({ Successful: acc.Successful + row.Successful, Failed: acc.Failed + row.Failed }), + { Successful: 0, Failed: 0 }, + ), + [chartData], + ); return ( - - -
- Success vs Failed Requests by Endpoint - + + {CATEGORIES.map((category, i) => ( + + + ))}
-
- - - -
+ } + > + + ); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.test.tsx index 81ca90d8cb9..8e26e63d9d9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.test.tsx @@ -2,6 +2,7 @@ import { screen } from "@testing-library/react"; import { describe, expect, it } from "vitest"; import { renderWithProviders } from "@/../tests/test-utils"; import { DailyData, MetricWithMetadata, SpendMetrics } from "@/components/UsagePage/types"; +import { STACKED_USAGE_PALETTE } from "@/components/shared/charts"; import EndpointUsageLineChart from "./EndpointUsageLineChart"; const spendMetrics = (apiRequests: number): SpendMetrics => ({ @@ -53,13 +54,13 @@ describe("EndpointUsageLineChart", () => { expect(screen.getByText("Endpoint Usage Trends")).toBeInTheDocument(); }); - it("renders one line per endpoint with the tremor palette strokes", () => { + it("renders one line per endpoint with the stacked usage palette strokes", () => { const { container } = renderWithProviders(); const curves = Array.from(container.querySelectorAll("path.recharts-line-curve")); expect(curves).toHaveLength(2); expect(new Set(curves.map((curve) => curve.getAttribute("stroke")))).toEqual( - new Set(["var(--color-blue-500, #3b82f6)", "var(--color-cyan-500, #06b6d4)"]), + new Set([STACKED_USAGE_PALETTE[0], STACKED_USAGE_PALETTE[1]]), ); }); @@ -87,7 +88,7 @@ describe("EndpointUsageLineChart", () => { expect(screen.getAllByText(/^\d,\d{3}$/).length).toBeGreaterThan(0); }); - it("draws smooth natural curves", () => { + it("draws smooth curves that never overshoot below zero (monotone)", () => { const { container } = renderWithProviders(); const path = container.querySelector("path.recharts-line-curve")?.getAttribute("d") ?? ""; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.tsx index 483a30b1639..ee2a50b5505 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageLineChart.tsx @@ -1,7 +1,8 @@ import { useMemo } from "react"; -import { LineChart, type ChartColor } from "@/components/shared/charts"; -import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card"; +import { ChartLine } from "lucide-react"; +import { LineChart, stackedUsageColor } from "@/components/shared/charts"; import { DailyData } from "@/components/UsagePage/types"; +import { Panel } from "../../overview/Primitives"; interface EndpointUsageLineChartProps { dailyData?: { results: DailyData[] }; @@ -58,41 +59,25 @@ export function EndpointUsageLineChart({ dailyData }: EndpointUsageLineChartProp return keys; }, [chartData]); - // Tremor color palette for multiple lines - const colors: readonly ChartColor[] = [ - "blue", - "cyan", - "indigo", - "violet", - "purple", - "fuchsia", - "pink", - "rose", - "red", - "orange", - ]; + // Same palette as the Overview / Model Leaderboard stacked chart + const colors = useMemo(() => categories.map((_, i) => stackedUsageColor(i)), [categories]); return ( - - - Endpoint Usage Trends - - - value.toLocaleString()} - showLegend={true} - showGridLines={true} - yAxisWidth={60} - connectNulls={true} - curveType="natural" - /> - - + + value.toLocaleString()} + showLegend={true} + showGridLines={true} + yAxisWidth={56} + connectNulls={true} + curveType="monotone" + /> + ); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageTable.tsx index 3d19d825c0c..2cccdfb47db 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageTable.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageTable.tsx @@ -1,5 +1,8 @@ import React from "react"; import type { ColumnDef } from "@tanstack/react-table"; +import { Route } from "lucide-react"; +import { cn } from "@/lib/cva.config"; +import { Panel } from "../../overview/Primitives"; import { Meter, MeterIndicator, MeterTrack } from "@/components/shared/Meter"; import { DataTable } from "@/components/shared/DataTable"; import { MoneyCell } from "@/components/shared/table_cells"; @@ -41,7 +44,7 @@ const EndpointUsageTable: React.FC = ({ endpointData }) { header: "Endpoint", accessorKey: "endpoint", - cell: ({ row }) => {row.original.endpoint}, + cell: ({ row }) => {row.original.endpoint}, }, { header: "Successful / Failed", @@ -54,18 +57,20 @@ const EndpointUsageTable: React.FC = ({ endpointData }) const totalPercentage = successPercentage + failurePercentage; return ( -
-
+
+
- 0 ? "bg-destructive" : undefined}> - + 0 ? "h-1 bg-destructive/70" : "h-1"}> +
-
- {record.successful_requests.toLocaleString()} +
+ {record.successful_requests.toLocaleString()} / - {record.failed_requests.toLocaleString()} + 0 ? "text-destructive" : "text-muted-foreground"}> + {record.failed_requests.toLocaleString()} +
); @@ -75,7 +80,7 @@ const EndpointUsageTable: React.FC = ({ endpointData }) header: "Total Request", accessorKey: "api_requests", meta: { numeric: true }, - cell: ({ row }) => row.original.api_requests.toLocaleString(), + cell: ({ row }) => {row.original.api_requests.toLocaleString()}, }, { header: "Success Rate", @@ -86,13 +91,10 @@ const EndpointUsageTable: React.FC = ({ endpointData }) const successRateStr = value.toFixed(2); return ( = 95 - ? "text-success font-medium" - : value >= 80 - ? "text-warning font-medium" - : "text-destructive font-medium" - } + className={cn( + "tabular-nums", + value >= 95 ? "text-foreground" : value >= 80 ? "text-warning" : "text-destructive", + )} > {successRateStr}% @@ -103,7 +105,7 @@ const EndpointUsageTable: React.FC = ({ endpointData }) header: "Total Tokens", accessorKey: "total_tokens", meta: { numeric: true }, - cell: ({ row }) => row.original.total_tokens.toLocaleString(), + cell: ({ row }) => {row.original.total_tokens.toLocaleString()}, }, { header: "Spend", @@ -114,13 +116,23 @@ const EndpointUsageTable: React.FC = ({ endpointData }) ]; return ( - row.key} - noDataMessage="No endpoint usage data" - size="compact" - /> + + {dataSource.length.toLocaleString()} {dataSource.length === 1 ? "endpoint" : "endpoints"} + + } + > + row.key} + noDataMessage="No endpoint usage data" + size="compact" + /> + ); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx index 23bd772729b..b546b2c0558 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.test.tsx @@ -563,7 +563,8 @@ describe("EntityUsage", () => { expect(spendElements.length).toBeGreaterThan(0); }); - expect(screen.getByText("1,000")).toBeInTheDocument(); // Total Requests + // Scoped to the active Cost tab: the keep-mounted Key Activity tab shows the same totals. + expect(within(screen.getByRole("tabpanel")).getByText("1,000")).toBeInTheDocument(); // Total Requests }); it("should render with team entity type and call team API", async () => { @@ -683,7 +684,7 @@ describe("EntityUsage", () => { fireEvent.click(keyActivityTab); }); - expect(screen.getAllByText("Activity Metrics")[1]).toBeInTheDocument(); + expect(within(screen.getByRole("tabpanel")).getByRole("heading", { name: "Overall Usage" })).toBeInTheDocument(); }); it("loads key pages separately from the aggregate using the current entity scope", async () => { @@ -806,9 +807,10 @@ describe("EntityUsage", () => { }); expect(await screen.findByText("Tag Spend Overview")).toBeInTheDocument(); - expect(await screen.findByText("$0.00")).toBeInTheDocument(); - expect(screen.getByText("Total Spend")).toBeInTheDocument(); - expect(screen.getAllByText("0")[0]).toBeInTheDocument(); + const costTab = screen.getByRole("tabpanel"); + expect(await within(costTab).findByText("$0.00")).toBeInTheDocument(); + expect(within(costTab).getByText("Total Spend")).toBeInTheDocument(); + expect(within(costTab).getAllByText("0")[0]).toBeInTheDocument(); }); it("should display Model Activity tab for non-agent entity types", async () => { @@ -1093,31 +1095,26 @@ describe("EntityUsage", () => { }); }); - it("renders daily spend bars, per-entity bars, and the provider donut with cyan fills and a $ center total", async () => { + it("renders the stacked daily spend chart, the per-entity table, and the provider share bar with a $ total", async () => { const { container } = render(); await waitFor(() => { expect(mockTagDailyActivityCall).toHaveBeenCalled(); }); + // The fixture carries no model breakdown, so the day's spend stacks as a single "Other" segment. await waitFor(() => { - expect(container.querySelectorAll("path.recharts-rectangle")).toHaveLength(2); + expect(container.querySelectorAll("path.recharts-rectangle")).toHaveLength(1); }); + expect(container.querySelector("path.recharts-rectangle")).toHaveAttribute("fill", "#94a3b8"); - const barFills = new Set( - Array.from(container.querySelectorAll("path.recharts-rectangle")).map((rect) => rect.getAttribute("fill")), - ); - expect(barFills).toEqual(new Set(["var(--color-cyan-500, #06b6d4)"])); + expect(screen.getAllByText("Jan 1").length).toBeGreaterThan(0); + expect(screen.getAllByText("Tag 1").length).toBeGreaterThan(0); - expect(screen.getAllByText("2025-01-01").length).toBeGreaterThan(0); - expect(screen.getAllByText("Tag 1").length).toBeGreaterThan(1); - - const sectors = container.querySelectorAll(".recharts-pie-sector path"); - expect(sectors).toHaveLength(1); - expect(sectors[0]).toHaveAttribute("fill", "var(--color-cyan-500, #06b6d4)"); - - const centerLabels = Array.from(container.querySelectorAll("text.fill-foreground")).map((text) => text.textContent); - expect(centerLabels).toContain("$100.50"); + const segments = screen.getAllByTestId("provider-share-segment"); + expect(segments).toHaveLength(1); + expect(segments[0]).toHaveStyle({ backgroundColor: "rgb(236, 72, 153)" }); + expect(screen.getByTestId("provider-spend-total")).toHaveTextContent("$100.50"); }); it("should label the chart with user_email metadata instead of the raw UUID (LIT-3889)", async () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx index 3f239250967..165fdfef703 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx @@ -1,32 +1,24 @@ import useTeams from "@/app/(dashboard)/hooks/useTeams"; -import { BarChart, DonutChart } from "@/components/shared/charts"; +import { StackedUsageChart, type StackedUsageScale } from "@/components/shared/charts"; import { DataTable } from "@/components/shared/DataTable"; -import { - getProviderSpend, - getTopAgents, - getTopAPIKeys, - getTopModels, - type ProviderSpendRow, -} from "./entityUsageAggregations"; +import { getProviderSpend, getTopAgents, getTopAPIKeys, getTopModels } from "./entityUsageAggregations"; import { buildCostBreakdownTiles, buildSummaryTiles, hasFlatCost, type SummaryTile } from "./entityUsageSummary"; import { MoneyCell } from "@/components/shared/table_cells"; -import { Card as ShadcnCard, CardContent, CardHeader, CardTitle } from "@/components/ui/card"; import { hasCapability, type Capability } from "@/utils/capabilities"; -import { formatNumberWithCommas } from "@/utils/dataUtils"; import type { DateRangePickerValue } from "@/components/shared/date_picker_types"; -import { ChevronDown, ChevronRight, Info } from "lucide-react"; +import { Bot, Boxes, ChevronDown, ChevronRight, ExternalLink, Info, KeyRound, Layers, Server } from "lucide-react"; import type { ColumnDef } from "@tanstack/react-table"; import { Alert, AlertDescription } from "@/components/shared/Alert"; import { ChartLoader } from "@/components/shared/chart_loader"; import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { Tooltip, TooltipContent, TooltipTrigger } from "@/components/ui/tooltip"; +import { cn } from "@/lib/cva.config"; import React, { type ReactNode, useCallback, useMemo, useState } from "react"; import TeamMultiSelect from "@/components/common_components/team_multi_select"; import UserDropdown from "@/components/common_components/UserDropdown"; import { ActivityMetrics, processActivityData } from "@/components/activity_metrics"; import { UsageExportHeader } from "@/components/EntityUsageExport"; import type { EntityType } from "@/components/EntityUsageExport/types"; -import { Logo } from "@/components/molecules/logo/Logo"; import { useAggregatedDailyActivity } from "../../hooks/useAggregatedDailyActivity"; import { ENTITY_API } from "./entityFetchFns"; import { @@ -35,28 +27,76 @@ import { type DailyActivityRequest, } from "@/components/UsagePage/dailyActivityApi"; import { keyDetailFromResponse, overallUsageMetrics } from "@/components/UsagePage/keyActivityData"; -import { EntityMetricWithMetadata } from "@/components/UsagePage/types"; -import { valueFormatterSpend } from "@/components/UsagePage/utils/value_formatters"; +import type { DailyData, EntityMetricWithMetadata } from "@/components/UsagePage/types"; import EndpointUsage from "../EndpointUsage/EndpointUsage"; import ModelViewToggle, { ModelViewType } from "../ModelViewToggle"; import TopKeyView from "@/components/UsagePage/components/EntityUsage/TopKeyView"; import KeyActivityPanel from "@/components/UsagePage/components/KeyActivityPanel"; +import { BreakdownControls, Leaderboard, useBreakdown, type BreakdownState } from "../overview/BreakdownChart"; +import { + bucketSeries, + bucketTotals, + dailyTotals, + formatMetricValue, + type Granularity, + type Series, + type UsageMetric, +} from "../overview/overviewData"; +import { Panel, Segmented, Sparkline, Stat } from "../overview/Primitives"; import TopModelView from "./TopModelView"; import TeamUserSpendCard from "./TeamUserSpendCard"; +import { ProviderSpendBreakdown } from "./SpendByProvider"; -interface EntityMetrics { - metrics: { - spend: number; - prompt_tokens: number; - completion_tokens: number; - cache_read_input_tokens: number; - cache_creation_input_tokens: number; - total_tokens: number; - successful_requests: number; - failed_requests: number; - api_requests: number; +const BRAND = "#2b3fd6"; +const FLAT_COST_SERIES = "Flat cost"; +const FLAT_COST_COLOR = "#8b5cf6"; +const METRIC_TITLE: Record = { spend: "Spend", tokens: "Tokens", requests: "Requests" }; +/** Summary tiles that carry a sparkline, keyed by tile title, valued by the dailyTotals field. */ +const TILE_TRENDS: Readonly> = { + "Total Spend": "spend", + "Total Cost": "spend", + "Total Requests": "requests", + "Total Tokens": "tokens", +}; +const GRANULARITY_OPTIONS = [ + { value: "day", label: "Daily" }, + { value: "week", label: "Weekly" }, +] as const satisfies readonly { value: Granularity; label: string }[]; +const SCALE_OPTIONS = [ + { value: "linear", label: "Linear" }, + { value: "log", label: "Log" }, +] as const satisfies readonly { value: StackedUsageScale; label: string }[]; +/** Quiet headers: muted, regular weight, so the rows carry the emphasis. */ +const QUIET_HEADER = { headerClassName: "font-normal" }; + +/** Stacks reserved-capacity flat cost on top of the per-model spend, so each bar is the day's full cost. */ +const withFlatCost = (series: Series, results: readonly DailyData[]): Series => { + const flatByDate = new Map(results.map((day) => [day.date, day.metrics.flat_cost ?? 0])); + return { + data: series.data.map((day) => ({ ...day, [FLAT_COST_SERIES]: flatByDate.get(day.date) ?? 0 })), + keys: [...series.keys, FLAT_COST_SERIES], + colors: [...series.colors, FLAT_COST_COLOR], }; - metadata: Record; +}; + +function ShareBar({ value, max }: { value: number; max: number }) { + return ( +