From 1238cfe90f1029bf33e31ab40a805e7ecabfa1ff Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 5 Oct 2026 01:11:52 -0700 Subject: [PATCH 01/19] feat(proxy): add LITELLM_FIPS_MODE startup gate with provider assertion and loud password migration failure (#42700) * feat(proxy): LITELLM_FIPS_MODE startup gate with provider assertion, TLS verify guard and loud password migration Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(proxy): match ssl_verify off detection to runtime str_to_bool semantics Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(proxy): drop tautological fips probe test and satisfy CodeQL return checks Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): inject the fake Prisma client through the module boundary instead of patching _setup_prisma_client Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yucheng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/common_utils/fips.py | 150 ++++++++++++++++++ litellm/proxy/proxy_server.py | 31 +++- tests/integration/_support/process.py | 54 +++++++ .../configuration/test_fips_mode_boot.py | 62 ++++++++ tests/unit/proxy/common_utils/test_fips.py | 106 +++++++++++++ ...test_proxy_server_endpoints_and_startup.py | 102 ++++++++++++ 6 files changed, 504 insertions(+), 1 deletion(-) create mode 100644 litellm/proxy/common_utils/fips.py create mode 100644 tests/integration/configuration/test_fips_mode_boot.py create mode 100644 tests/unit/proxy/common_utils/test_fips.py diff --git a/litellm/proxy/common_utils/fips.py b/litellm/proxy/common_utils/fips.py new file mode 100644 index 00000000000..07bd4d2462c --- /dev/null +++ b/litellm/proxy/common_utils/fips.py @@ -0,0 +1,150 @@ +import hashlib +import os +from collections.abc import Callable +from dataclasses import dataclass +from typing import Final + +from typing_extensions import assert_never + +from litellm.secret_managers.main import str_to_bool + +FIPS_MODE_ENV_VAR: Final = "LITELLM_FIPS_MODE" +SSL_VERIFY_ENV_VAR: Final = "SSL_VERIFY" +SSL_VERIFY_SETTING: Final = "litellm_settings.ssl_verify" +REFUSAL_PREFIX: Final = "LiteLLM proxy refused to start" + +_TRUE_VALUES: Final = frozenset({"true", "1", "yes", "on"}) +_FALSE_VALUES: Final = frozenset({"false", "0", "no", "off", ""}) + + +@dataclass(frozen=True, slots=True) +class FipsModeOff: + pass + + +@dataclass(frozen=True, slots=True) +class FipsModeOn: + pass + + +@dataclass(frozen=True, slots=True) +class MalformedFipsMode: + value: str + + +FipsModeSetting = FipsModeOff | FipsModeOn | MalformedFipsMode + + +@dataclass(frozen=True, slots=True) +class ProviderDoesNotEnforceFips: + pass + + +@dataclass(frozen=True, slots=True) +class TlsVerificationDisabled: + sources: tuple[str, ...] + + +FipsBootRefusal = MalformedFipsMode | ProviderDoesNotEnforceFips | TlsVerificationDisabled +FipsBootVerdict = FipsModeOff | FipsModeOn | FipsBootRefusal + + +class FipsModeError(Exception): + pass + + +def parse_fips_mode(raw: str | None) -> FipsModeSetting: + if raw is None: + return FipsModeOff() + normalized: Final = raw.strip().lower() + if normalized in _TRUE_VALUES: + return FipsModeOn() + if normalized in _FALSE_VALUES: + return FipsModeOff() + return MalformedFipsMode(value=raw) + + +def is_fips_mode(environ: Callable[[str], str | None] = os.environ.get) -> bool: + return isinstance(parse_fips_mode(environ(FIPS_MODE_ENV_VAR)), FipsModeOn) + + +def openssl_enforces_fips() -> bool: + """MD5 is not an approved digest, so an enforcing FIPS provider refuses it even when asked for security use.""" + try: + hashlib.md5(b"", usedforsecurity=True) + except ValueError: + return True + return False + + +def fips_boot_verdict( + *, + raw_fips_mode: str | None, + provider_enforces_fips: Callable[[], bool], + ssl_verify_environment: str | None, + ssl_verify_setting: object, +) -> FipsBootVerdict: + setting: Final = parse_fips_mode(raw_fips_mode) + match setting: + case FipsModeOff() | MalformedFipsMode(): + return setting + case FipsModeOn(): + pass + case _: + assert_never(setting) + disabled: Final = tuple( + source + for source, off in ( + (SSL_VERIFY_ENV_VAR, _is_off(ssl_verify_environment)), + (SSL_VERIFY_SETTING, _is_off(ssl_verify_setting)), + ) + if off + ) + if disabled: + return TlsVerificationDisabled(sources=disabled) + if not provider_enforces_fips(): + return ProviderDoesNotEnforceFips() + return setting + + +def enforce_fips_boot_verdict(verdict: FipsBootVerdict, announce: Callable[[str], object]) -> None: + match verdict: + case FipsModeOff() | FipsModeOn(): + return + case MalformedFipsMode() | ProviderDoesNotEnforceFips() | TlsVerificationDisabled(): + message: Final = render_refusal(verdict) + announce(f"\n{message}\n\n") + raise FipsModeError(message) + case _: + assert_never(verdict) + + +def render_refusal(refusal: FipsBootRefusal) -> str: + match refusal: + case MalformedFipsMode(): + return ( + f"{REFUSAL_PREFIX}: {FIPS_MODE_ENV_VAR}={refusal.value} is not a boolean.\n" + f"Set {FIPS_MODE_ENV_VAR} to true or false, or unset it." + ) + case ProviderDoesNotEnforceFips(): + return ( + f"{REFUSAL_PREFIX}: {FIPS_MODE_ENV_VAR} is on but this Python does not enforce FIPS.\n" + "Its OpenSSL still allows non-approved algorithms (MD5 succeeded), so passwords and keys would be\n" + "protected with algorithms the FIPS 140-3 policy forbids. Run the proxy from a FIPS image whose\n" + f"OpenSSL FIPS provider is enabled, or unset {FIPS_MODE_ENV_VAR} on a non-FIPS runtime." + ) + case TlsVerificationDisabled(): + return ( + f"{REFUSAL_PREFIX}: {FIPS_MODE_ENV_VAR} is on but TLS certificate verification is disabled by " + f"{' and '.join(refusal.sources)}.\nFIPS deployments must verify upstream certificates, so remove the " + "override or point ssl_verify at a CA bundle instead." + ) + return assert_never(refusal) + + +def _is_off(value: object) -> bool: + if isinstance(value, bool): + return value is False + if isinstance(value, str): + return str_to_bool(value) is False + return False diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f40e541dd86..af487cb17a3 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -422,6 +422,14 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( encrypt_value_helper, ) from litellm.proxy.common_utils.error_body_call_id import JSON_OBJECT, error_body_call_id, with_call_id +from litellm.proxy.common_utils.fips import ( + FIPS_MODE_ENV_VAR, + SSL_VERIFY_ENV_VAR, + enforce_fips_boot_verdict, + fips_boot_verdict, + is_fips_mode, + openssl_enforces_fips, +) from litellm.proxy.common_utils.healthy_model_filter import ( get_hidden_unhealthy_model_names, is_healthy_only_listing_default, @@ -1329,6 +1337,16 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState if isinstance(worker_config, dict): await initialize_from_worker_config(worker_config) + enforce_fips_boot_verdict( + fips_boot_verdict( + raw_fips_mode=os.getenv(FIPS_MODE_ENV_VAR), + provider_enforces_fips=openssl_enforces_fips, + ssl_verify_environment=os.getenv(SSL_VERIFY_ENV_VAR), + ssl_verify_setting=litellm.ssl_verify, + ), + announce=announce_on_stderr_at_exit, + ) + enforce_master_key_boot_verdict( await with_stored_secrets_counted( master_key_boot_verdict( @@ -1368,10 +1386,21 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState try: result: Final = await migrate_passwords_to_scrypt_async(prisma_client) verbose_proxy_logger.info("Password migration: %s", result) + except ValueError as e: + verbose_proxy_logger.error( + "Password migration failed, so plaintext passwords stay unhashed in the database: %s. " + "This is what an OpenSSL FIPS provider reports when the hashing algorithm is not approved.", + e, + ) + if is_fips_mode(): + raise except Exception as e: verbose_proxy_logger.warning("Password migration skipped: %s", e) - asyncio.create_task(_run_pw_migration()) + if is_fips_mode(): + await _run_pw_migration() + else: + asyncio.create_task(_run_pw_migration()) async def _run_agent_grant_id_migration() -> None: from litellm.proxy.agent_endpoints.agent_registry import ( diff --git a/tests/integration/_support/process.py b/tests/integration/_support/process.py index 3d76c7881e1..3fa7f4b0333 100644 --- a/tests/integration/_support/process.py +++ b/tests/integration/_support/process.py @@ -230,6 +230,60 @@ def owned_proxy_process( _stop(process) +def _is_ready(client: httpx.Client) -> bool: + try: + return client.get("/health/readiness", timeout=2).status_code == 200 + except httpx.TransportError: + return False + + +def refused_boot_log( + gateway: Gateway, + directory: Path, + overrides: Mapping[str, str], + *, + config: Path | None = None, +) -> str: + """Start the proxy and return its log once it exits non-zero instead of becoming ready.""" + root: Final = Path(os.environ.get("INTEGRATION_PROXY_ROOT") or Path(__file__).resolve().parents[3]) + environment: Final = { + **os.environ, + **proxy_database_environment(), + "LITELLM_MASTER_KEY": gateway.key, + "LITELLM_SALT_KEY": os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt"), + "STORE_MODEL_IN_DB": "True", + **overrides, + } + output: Final = Path(os.environ.get("INTEGRATION_RESULTS_DIR", str(directory))) + output.mkdir(parents=True, exist_ok=True) + command: Final = ( + sys.executable, + "-m", + "integration._support.proxy", + "--config", + str(config or "tests/integration/proxy_config.yaml"), + "--host", + "127.0.0.1", + "--num_workers", + "1", + *DB_PUSH, + ) + launch: Final = _launch(command, root, environment, output) + try: + with httpx.Client(base_url=f"http://127.0.0.1:{launch.port}", timeout=15, trust_env=False) as client: + deadline: Final = time.monotonic() + 70 + while launch.process.poll() is None: + assert not _is_ready(client), ( + f"Proxy became ready instead of refusing to boot:\n{launch.log.read_text()}" + ) + assert time.monotonic() < deadline, "Proxy neither exited nor became ready within the deadline" + time.sleep(0.1) + assert launch.process.returncode != 0, f"Proxy exited 0 instead of refusing to boot:\n{launch.log.read_text()}" + return launch.log.read_text() + finally: + _stop(launch.process) + + _UPSTREAM_READY_SECONDS: Final = 60 diff --git a/tests/integration/configuration/test_fips_mode_boot.py b/tests/integration/configuration/test_fips_mode_boot.py new file mode 100644 index 00000000000..80904c600e1 --- /dev/null +++ b/tests/integration/configuration/test_fips_mode_boot.py @@ -0,0 +1,62 @@ +"""LITELLM_FIPS_MODE is a boot gate: the proxy refuses to serve unless the process really enforces FIPS. + +Every leg launches the real proxy binary against the suite's Postgres and asserts on what an operator sees: +exit status and the refusal text in the log. Nothing is patched inside the proxy. +""" + +import hashlib +from pathlib import Path +from typing import Final + +import pytest +import yaml + +from tests.integration._support.client import Gateway +from tests.integration._support.process import owned_proxy, refused_boot_log + +REFUSAL: Final = "LiteLLM proxy refused to start" + + +def _this_python_enforces_fips() -> bool: + try: + hashlib.md5(b"probe", usedforsecurity=True) + except ValueError: + return True + return False + + +def test_fips_mode_refuses_to_serve_when_this_python_does_not_enforce_fips(gateway: Gateway, tmp_path: Path) -> None: + if _this_python_enforces_fips(): + pytest.skip("Runner OpenSSL enforces FIPS, so this leg cannot observe the non-enforcing refusal") + log: Final = refused_boot_log(gateway, tmp_path, {"LITELLM_FIPS_MODE": "true"}) + assert REFUSAL in log, log + assert "LITELLM_FIPS_MODE" in log and "does not enforce FIPS" in log, log + + +@pytest.mark.parametrize("source", ("environment", "config")) +def test_fips_mode_refuses_to_serve_with_tls_verification_disabled( + gateway: Gateway, tmp_path: Path, source: str +) -> None: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + path: Final = tmp_path / "ssl_verify_off.yaml" + settings: Final = {**config.get("litellm_settings", {}), "ssl_verify": False} + path.write_text(yaml.safe_dump({**config, "litellm_settings": settings})) + log: Final = ( + refused_boot_log(gateway, tmp_path, {"LITELLM_FIPS_MODE": "true", "SSL_VERIFY": "false"}) + if source == "environment" + else refused_boot_log(gateway, tmp_path, {"LITELLM_FIPS_MODE": "true"}, config=path) + ) + assert REFUSAL in log, log + assert "TLS certificate verification is disabled" in log, log + assert ("SSL_VERIFY" if source == "environment" else "litellm_settings.ssl_verify") in log, log + + +def test_fips_mode_refuses_to_serve_on_a_value_that_is_not_a_boolean(gateway: Gateway, tmp_path: Path) -> None: + log: Final = refused_boot_log(gateway, tmp_path, {"LITELLM_FIPS_MODE": "enforced"}) + assert REFUSAL in log, log + assert "LITELLM_FIPS_MODE=enforced" in log and "true or false" in log, log + + +def test_fips_mode_off_serves_even_with_tls_verification_disabled(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy(gateway, tmp_path, {"LITELLM_FIPS_MODE": "false", "SSL_VERIFY": "false"}) as candidate: + assert candidate.client.get("/health/readiness").status_code == 200 diff --git a/tests/unit/proxy/common_utils/test_fips.py b/tests/unit/proxy/common_utils/test_fips.py new file mode 100644 index 00000000000..0e8138856c0 --- /dev/null +++ b/tests/unit/proxy/common_utils/test_fips.py @@ -0,0 +1,106 @@ +import pytest + +from litellm.proxy.common_utils.fips import ( + FipsModeError, + FipsModeOff, + FipsModeOn, + MalformedFipsMode, + ProviderDoesNotEnforceFips, + TlsVerificationDisabled, + enforce_fips_boot_verdict, + fips_boot_verdict, + is_fips_mode, + parse_fips_mode, +) + + +def _verdict( + raw: str | None, + *, + enforcing: bool = True, + ssl_env: str | None = None, + ssl_setting: object = True, +): + return fips_boot_verdict( + raw_fips_mode=raw, + provider_enforces_fips=lambda: enforcing, + ssl_verify_environment=ssl_env, + ssl_verify_setting=ssl_setting, + ) + + +@pytest.mark.parametrize("raw", [None, "false", "0", "no", "off", "", " False "]) +def test_unset_and_false_spellings_leave_fips_mode_off(raw): + assert parse_fips_mode(raw) == FipsModeOff() + assert is_fips_mode({"LITELLM_FIPS_MODE": raw}.get) is False + + +@pytest.mark.parametrize("raw", ["true", "1", "yes", "on", " TRUE "]) +def test_true_spellings_turn_fips_mode_on(raw): + assert parse_fips_mode(raw) == FipsModeOn() + assert is_fips_mode({"LITELLM_FIPS_MODE": raw}.get) is True + + +@pytest.mark.parametrize("raw", ["enforced", "2", "strict", "yes please"]) +def test_anything_else_is_malformed_and_refused_with_the_offending_value(raw): + assert parse_fips_mode(raw) == MalformedFipsMode(value=raw) + with pytest.raises(FipsModeError) as refused: + enforce_fips_boot_verdict(_verdict(raw, enforcing=False), announce=lambda _: None) + assert f"LITELLM_FIPS_MODE={raw} is not a boolean" in str(refused.value) + assert "true or false" in str(refused.value) + + +def test_off_never_consults_the_provider_or_tls_settings(): + def explode() -> bool: + raise AssertionError("provider probe must not run while FIPS mode is off") + + verdict = fips_boot_verdict( + raw_fips_mode=None, provider_enforces_fips=explode, ssl_verify_environment="false", ssl_verify_setting=False + ) + assert verdict == FipsModeOff() + enforce_fips_boot_verdict(verdict, announce=lambda _: pytest.fail("nothing to announce when off")) + + +def test_on_with_an_enforcing_provider_and_verified_tls_boots(): + verdict = _verdict("true", enforcing=True, ssl_env="true", ssl_setting="/etc/ssl/certs/ca.pem") + assert verdict == FipsModeOn() + enforce_fips_boot_verdict(verdict, announce=lambda _: pytest.fail("nothing to announce when on")) + + +def test_on_with_a_non_enforcing_provider_is_refused_and_names_the_fix(): + announced = [] + with pytest.raises(FipsModeError) as refused: + enforce_fips_boot_verdict(_verdict("true", enforcing=False), announce=announced.append) + message = str(refused.value) + assert message.startswith("LiteLLM proxy refused to start") + assert "LITELLM_FIPS_MODE is on but this Python does not enforce FIPS" in message + assert "FIPS image" in message + assert announced == [f"\n{message}\n\n"] + + +@pytest.mark.parametrize( + "ssl_env, ssl_setting, sources", + [ + ("false", True, ("SSL_VERIFY",)), + (" FALSE ", True, ("SSL_VERIFY",)), + (None, False, ("litellm_settings.ssl_verify",)), + (None, "False", ("litellm_settings.ssl_verify",)), + ("false", False, ("SSL_VERIFY", "litellm_settings.ssl_verify")), + ], +) +def test_disabled_tls_verification_is_refused_naming_every_source(ssl_env, ssl_setting, sources): + verdict = _verdict("true", enforcing=True, ssl_env=ssl_env, ssl_setting=ssl_setting) + assert verdict == TlsVerificationDisabled(sources=sources) + with pytest.raises(FipsModeError) as refused: + enforce_fips_boot_verdict(verdict, announce=lambda _: None) + assert "TLS certificate verification is disabled by " + " and ".join(sources) in str(refused.value) + + +@pytest.mark.parametrize("ssl_setting", [True, "true", "/etc/ssl/certs/ca.pem", None, "", "0", "no"]) +def test_verified_or_custom_bundle_tls_settings_are_not_treated_as_disabled(ssl_setting): + assert _verdict("true", enforcing=True, ssl_setting=ssl_setting) == FipsModeOn() + + +def test_disabled_tls_is_reported_before_the_provider_so_operators_see_config_mistakes_first(): + assert _verdict("true", enforcing=False, ssl_env="false") == TlsVerificationDisabled(sources=("SSL_VERIFY",)) + assert _verdict("true", enforcing=False) == ProviderDoesNotEnforceFips() diff --git a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py index 42a9dc441dd..2397c6480d5 100644 --- a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py +++ b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py @@ -1771,6 +1771,108 @@ async def test_proxy_startup_refuses_an_unsafe_master_key_even_when_the_database assert ("could not be checked" in announced[0]) == key_can_have_encrypted_the_database +@pytest.mark.asyncio +async def test_proxy_startup_refuses_fips_mode_when_this_python_does_not_enforce_fips(monkeypatch, tmp_path): + from fastapi import FastAPI + + from litellm.proxy.common_utils.fips import FipsModeError + from litellm.proxy.proxy_server import proxy_startup_event + + _, announced = _boot_with_general_settings(monkeypatch, tmp_path, {"master_key": "sk-a-safe-master-key"}) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.openssl_enforces_fips", lambda: False) + monkeypatch.setenv("LITELLM_FIPS_MODE", "true") + + with pytest.raises(FipsModeError): + async with proxy_startup_event(FastAPI()): + pass + + assert len(announced) == 1 + assert "does not enforce FIPS" in announced[0] + + +@pytest.mark.asyncio +async def test_proxy_startup_refuses_fips_mode_when_the_config_disables_tls_verification(monkeypatch, tmp_path): + import yaml + from fastapi import FastAPI + + from litellm.proxy.common_utils.fips import FipsModeError + from litellm.proxy.proxy_server import proxy_startup_event + + config_path, announced = _boot_with_general_settings(monkeypatch, tmp_path, {"master_key": "sk-a-safe-master-key"}) + config_path.write_text( + yaml.dump({"general_settings": {"master_key": "sk-a-safe-master-key"}, "litellm_settings": {"ssl_verify": False}}) + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.openssl_enforces_fips", lambda: True) + monkeypatch.setattr(litellm, "ssl_verify", True) + monkeypatch.setenv("LITELLM_FIPS_MODE", "true") + + with pytest.raises(FipsModeError): + async with proxy_startup_event(FastAPI()): + pass + + assert "TLS certificate verification is disabled by litellm_settings.ssl_verify" in announced[0] + + +class _PrismaClientWhoseUserTableCannotHash: + class _Table: + async def find_many(self, where): + raise ValueError("[digital envelope routines] unsupported") + + class _Db: + litellm_usertable = None + + def __init__(self, database_url, proxy_logging_obj): + self.db = self._Db() + self.db.litellm_usertable = self._Table() + self.writer_db = self.db + + async def connect(self): + pass + + async def disconnect(self): + pass + + def start_view_setup_task(self): + pass + + async def check_view_exists(self): + pass + + async def health_check(self): + pass + + +@pytest.mark.asyncio +@pytest.mark.parametrize("fips_mode", ["true", "false"]) +async def test_proxy_startup_surfaces_a_password_migration_crypto_failure(monkeypatch, tmp_path, caplog, fips_mode): + from fastapi import FastAPI + + from litellm.proxy.proxy_server import proxy_startup_event + + _boot_with_general_settings(monkeypatch, tmp_path, {"master_key": "sk-a-safe-master-key"}) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setenv("DATABASE_URL", "postgresql://nobody:nothing@127.0.0.1:1/unreachable") + monkeypatch.setattr("litellm.proxy.proxy_server.PrismaClient", _PrismaClientWhoseUserTableCannotHash) + monkeypatch.setattr("litellm.proxy.proxy_server.openssl_enforces_fips", lambda: True) + monkeypatch.setenv("LITELLM_FIPS_MODE", fips_mode) + + with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"): + if fips_mode == "true": + with pytest.raises(ValueError, match="digital envelope routines"): + async with proxy_startup_event(FastAPI()): + pass + else: + async with proxy_startup_event(FastAPI()): + await asyncio.sleep(0) + + failures = [r.getMessage() for r in caplog.records if "Password migration failed" in r.getMessage()] + assert len(failures) == 1 + assert "plaintext passwords stay unhashed" in failures[0] + assert "digital envelope routines" in failures[0] + + class _DatabaseWithOneStoredCredential: def __init__(self, ciphertext): self._ciphertext = ciphertext From 7a7d27c550fa4edd353d1dcbfb9ab8c13ab096b7 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 5 Oct 2026 01:44:43 -0700 Subject: [PATCH 02/19] fix(guardrails): run the end-of-stream post_call scan when the client disconnects mid-stream (#43839) * fix(guardrails): run end-of-stream post_call scan when the client disconnects mid-stream Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): close the guardrail stream chain in async_data_generator on client disconnect Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): leave the raw upstream response to the shielded finalizer on client disconnect Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): keep disconnect cleanup going when a streaming callback cleanup raises Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): assert the refund through a recorder instead of the mock Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(guardrails): inspect tool calls released before a disconnect under incremental_diff and record a failed scan marker The incremental_diff transform stream now scans tool calls it already released when the client disconnects, and a disconnect scan whose translation raises after the guardrail recorded success also records guardrail_failed_to_respond Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): pin that a text-only disconnect scan is not handed a tool_calls finish Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): cover disconnect scans on every streaming endpoint and client, plus outage, worker-kill and cache-hit cells Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): prove the cache-hit twin is served from the cache Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(guardrails): type the disconnect-close streams so basedpyright stops reporting unknown arguments * fix(guardrails): give the guardrail metadata cast a reason so the type discipline gate accepts it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(guardrails): scan released Messages and Responses tool calls on disconnect Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(guardrails): format the disconnect scan unit tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * chore(guardrails): drop mutable-ok markers that no longer suppress a rule Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(guardrails): scan released Responses output after a finished item and end only in-flight Chat choices Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): type the disconnect scan test helpers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(guardrails): type request_data in the disconnect scan helpers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): type request_data in the disconnect scan test doubles Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): pin that chat streams with no tool call in flight end as released Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(guardrails): scan only released chunks on disconnect and skip it once a block owns the verdict The disconnect scan now uses the chunks actually yielded to the client, copies them before scanning, skips when a mid-stream block or HTTP error already settled the verdict, and the iterator wrapper only closes hooks that are async generators so plain async iterator hooks keep working Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(guardrails): pin that a delivered guardrail error or final chunk settles the disconnect verdict Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(guardrails): close any hook iterator that exposes aclose when the stream ends early Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): accept a synchronous aclose on custom streaming hook iterators Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): swallow callback aclose errors at end of stream A custom callback whose async_post_call_streaming_iterator_hook returns a non-generator async iterator with a raising aclose() failed the finished stream: content plus usage reached the client and then the stream surfaced an error SSE with no [DONE], or aborted a post_call pipeline's buffering loop into a 500 with an empty body. Wrap the aclose invocation in _wrap_streaming_iterator_with_enrichment in try/except and log a warning naming the callback and the cleanup error, matching close_guarded_stream and _close_guarded_layers. Iteration-time hook exceptions still propagate. * fix(proxy): log only the error type when a callback aclose raises Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yucheng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../chat/guardrail_translation/handler.py | 11 + .../guardrail_translation/base_translation.py | 5 + .../chat/guardrail_translation/handler.py | 29 +- .../guardrail_translation/handler.py | 66 ++ litellm/proxy/common_request_processing.py | 24 +- .../guardrail_hooks/bedrock_guardrails.py | 20 +- .../unified_guardrail/unified_guardrail.py | 521 +++++++----- litellm/proxy/proxy_server.py | 27 +- litellm/proxy/utils.py | 40 +- .../observability/test_guardrail_effects.py | 749 +++++++++++++++++- .../test_anthropic_guardrail_handler.py | 22 + .../test_openai_guardrail_handler.py | 29 + ...test_openai_responses_guardrail_handler.py | 83 ++ .../test_bedrock_guardrails.py | 35 + .../test_unified_guardrail.py | 563 ++++++++++++- ...async_post_call_streaming_iterator_hook.py | 340 +++++++- .../proxy/test_common_request_processing.py | 67 ++ ...test_proxy_server_endpoints_and_startup.py | 63 ++ .../proxy_logging/test_guardrail_pipeline.py | 62 +- 19 files changed, 2527 insertions(+), 229 deletions(-) diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 806240c9749..578a4e47dfd 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -224,6 +224,11 @@ def _write_back_message_text(message: _WritableMessage, target: MessageTextTarge _TOOL_USE_INPUT_ADAPTER: Final = TypeAdapter(dict[str, object]) +_RELEASED_TOOL_USE_STOP: Final = ( + b"event: message_delta\n" + b'data: {"type": "message_delta", "delta": {"stop_reason": "tool_use", "stop_sequence": null}, ' + b'"usage": {"output_tokens": 0}}\n\n' +) def _rewritten_tool_use_input(arguments: str) -> Mapping[str, object] | None: @@ -1573,6 +1578,12 @@ class AnthropicMessagesHandler(BaseTranslation): tool_calls_in_flight=bool(tool_use_fingerprints) and not stream_ended, ) + def released_stream_as_ended(self, responses_so_far: Sequence[object]) -> tuple[object, ...]: + released_key: Final = self.get_streaming_scan_key(responses_so_far) + if released_key is None or not released_key.tool_calls_in_flight: + return tuple(responses_so_far) + return (*responses_so_far, _RELEASED_TOOL_USE_STOP) + @classmethod def _streamed_tool_use_fingerprints(cls, responses_so_far: Sequence[object]) -> tuple[str, ...]: return tuple( diff --git a/litellm/llms/base_llm/guardrail_translation/base_translation.py b/litellm/llms/base_llm/guardrail_translation/base_translation.py index 89ad67f0485..70b8291d32c 100644 --- a/litellm/llms/base_llm/guardrail_translation/base_translation.py +++ b/litellm/llms/base_llm/guardrail_translation/base_translation.py @@ -203,6 +203,11 @@ class BaseTranslation(ABC): def get_streaming_scan_key(self, responses_so_far: Sequence[object]) -> StreamingScanKey | None: return None + def released_stream_as_ended(self, responses_so_far: Sequence[object]) -> tuple[object, ...]: + """The chunks a client left the stream with, closed the way this endpoint ends a stream, so the + end-of-stream scan also inspects tool calls the stream never finished""" + return tuple(responses_so_far) + def build_block_sse_chunks( self, exc: "ModifyResponseException", diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index aa175733582..aee439862db 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -17,7 +17,7 @@ This pattern can be replicated for other message formats (e.g., Anthropic). import json import time import uuid -from collections.abc import Mapping, Sequence +from collections.abc import Iterator, Mapping, Sequence from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Union, cast @@ -54,6 +54,7 @@ from litellm.types.utils import ( ChatCompletionDeltaToolCall, ChatCompletionMessageToolCall, Choices, + Delta, GenericGuardrailAPIInputs, ModelResponse, ModelResponseStream, @@ -837,6 +838,18 @@ class OpenAIChatCompletionsHandler(BaseTranslation): tool_calls_in_flight=bool(tool_call_fingerprints) and not stream_ended, ) + def released_stream_as_ended(self, responses_so_far: Sequence[object]) -> tuple[object, ...]: + released_key: Final = self.get_streaming_scan_key(responses_so_far) + if released_key is None or not released_key.tool_calls_in_flight: + return tuple(responses_so_far) + terminator: Final = ModelResponseStream( + choices=[ + StreamingChoices(index=index, delta=Delta(), finish_reason="tool_calls") + for index in _choice_indices_with_tool_calls(responses_so_far) + ] + ) + return (*responses_so_far, terminator) + @staticmethod def _streamed_tool_call_fingerprints(responses_so_far: Sequence[object]) -> tuple[str, ...]: return tuple( @@ -1388,6 +1401,20 @@ def _streamed_delta_tool_calls(delta: object) -> tuple[object, ...]: return stream_item_items(delta, "tool_calls") + legacy +def _released_choices(responses_so_far: Sequence[object]) -> Iterator[object]: + for chunk in responses_so_far: + yield from _stream_chunk_choices(chunk) + + +def _choice_indices_with_tool_calls(responses_so_far: Sequence[object]) -> tuple[int, ...]: + indices: Final = ( + index if isinstance(index := stream_item_field(choice, "index"), int) else 0 + for choice in _released_choices(responses_so_far) + if _streamed_delta_tool_calls(stream_item_field(choice, "delta")) + ) + return tuple(dict.fromkeys(indices)) + + def _blocked_stream_identity( exc: "ModifyResponseException", responses_so_far: Sequence[object] ) -> tuple[str, int, str]: diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index 90cdef87ec7..ad68924ce20 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -293,6 +293,51 @@ def _is_tool_call_output_item(item: object) -> bool: return _tool_call_output_item_mapping(item) is not None +def _released_tool_call_payload(responses_so_far: Sequence[object], item_id: object) -> str | None: + events: Final = tuple(event for event in responses_so_far if stream_item_field(event, "item_id") == item_id) + finished: Final = tuple( + payload + for event in events + if isinstance(event_type := stream_item_field(event, "type"), str) + and event_type in _TOOL_CALL_PAYLOAD_DONE_EVENT_FIELDS + and isinstance(payload := stream_item_field(event, _TOOL_CALL_PAYLOAD_DONE_EVENT_FIELDS[event_type]), str) + ) + if finished: + return finished[-1] + deltas: Final = tuple( + delta + for event in events + if stream_item_field(event, "type") in _TOOL_CALL_PAYLOAD_DELTA_EVENT_TYPES + and isinstance(delta := stream_item_field(event, "delta"), str) + ) + return "".join(deltas) if deltas else None + + +def _with_released_payload(item: Mapping[str, object], responses_so_far: Sequence[object]) -> Mapping[str, object]: + field: Final = _TOOL_CALL_PAYLOAD_FIELDS[str(item.get("type"))] + payload: Final = _released_tool_call_payload(responses_so_far, item.get("id")) + return item if payload is None else {**item, field: payload} + + +def _released_message_item(text: str) -> Mapping[str, object]: + content: Final = [{"type": "output_text", "text": text}] + return {"type": "message", "role": "assistant", "content": content} + + +def _released_tool_call_items(responses_so_far: Sequence[object]) -> tuple[Mapping[str, object], ...]: + announced: Final = tuple( + item + for item in ( + _tool_call_output_item_mapping(stream_item_field(event, "item")) + for event in responses_so_far + if stream_item_field(event, "type") in _OUTPUT_ITEM_EVENT_TYPES + ) + if item is not None + ) + latest_by_id: Final = MappingProxyType({item.get("id"): item for item in announced}) + return tuple(_with_released_payload(item, responses_so_far) for item in latest_by_id.values()) + + def _last_message_role(messages: Sequence[object]) -> str | None: if not messages: return None @@ -1324,6 +1369,27 @@ class OpenAIResponsesHandler(BaseTranslation): tool_calls_in_flight=self._has_streamed_tool_call_events(responses_so_far), ) + def released_stream_as_ended(self, responses_so_far: Sequence[object]) -> tuple[object, ...]: + if self._check_streaming_has_ended(responses_so_far): + return tuple(responses_so_far) + ends_on_finished_item: Final = ( + bool(responses_so_far) + and stream_item_field(responses_so_far[-1], "type") == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE.value + ) + if not ends_on_finished_item and not self._has_streamed_tool_call_events(responses_so_far): + return tuple(responses_so_far) + text_events: Final = tuple( + event for event in responses_so_far if stream_item_field(event, "type") in _OUTPUT_TEXT_EVENT_TYPES + ) + released_text: Final = self.get_streaming_string_so_far(text_events) + message_items: Final = (_released_message_item(released_text),) if released_text else () + tool_items: Final = _released_tool_call_items(responses_so_far) + output: Final = [*message_items, *tool_items] + response: Final = {"status": "incomplete", "output": output} + incomplete: Final = ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE.value + envelope: Final = {"type": incomplete, "response": response} + return (*responses_so_far, envelope) + @staticmethod def _has_streamed_tool_call_events(responses_so_far: Sequence[object]) -> bool: return any( diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index c85169f0ba5..687df9b0348 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -272,6 +272,18 @@ def _withheld_provider_output(response: object) -> bool: return getattr(response, "has_buffered_provider_output", False) is True +async def close_guarded_stream(stream: object) -> None: + if not isinstance(stream, AsyncGenerator): + return + with anyio.CancelScope(shield=True): + try: + await stream.aclose() + except Exception as e: # noqa: BLE001 # a failing callback cleanup must not skip the refund and finalizer + verbose_proxy_logger.warning( + "Closing the guarded stream after a client disconnect raised %s", type(e).__name__ + ) + + def resolve_litellm_call_id(client_call_id: str | None) -> str: if client_call_id is not None and 0 < len(client_call_id) <= MAX_LITELLM_CALL_ID_LENGTH: return client_call_id @@ -3909,13 +3921,14 @@ class ProxyBaseLLMRequestProcessing: client_disconnected = False delivered_chunk = False recent_tail = SSE_STREAM_START_TAIL # rebind-ok: rolling window over the yielded bytes + guarded_stream: Final[AsyncGenerator[object, None]] = proxy_logging_obj.async_post_call_streaming_iterator_hook( + user_api_key_dict=user_api_key_dict, + response=response, + request_data=request_data, + ) try: str_so_far = "" - async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook( - user_api_key_dict=user_api_key_dict, - response=response, - request_data=request_data, - ): + async for chunk in guarded_stream: # ``.format(chunk)`` was previously evaluated for every chunk # regardless of log level; gate it behind the level check. if debug_enabled: @@ -3971,6 +3984,7 @@ class ProxyBaseLLMRequestProcessing: # Starlette closes on disconnect, so the nested iterator hook (which # only sees GeneratorExit on GC) cannot own the refund. client_disconnected = not stream_completed + await close_guarded_stream(guarded_stream) if not delivered_chunk and not _withheld_provider_output(response): from litellm.proxy.spend_tracking.budget_reservation import ( release_budget_reservation_on_cancel, diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 620b24df95d..f900d14bdc1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -10,6 +10,7 @@ import sys sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path import asyncio +import contextlib import copy import json import re @@ -2747,14 +2748,17 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): UnifiedLLMGuardrails, ) - async for streamed_chunk in UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( - user_api_key_dict=user_api_key_dict, - response=response, - request_data=request_data, - guardrail_to_apply=self, - buffer_until_moderated_default=False, - ): - yield streamed_chunk + async with contextlib.aclosing( + UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( + user_api_key_dict=user_api_key_dict, + response=response, + request_data=request_data, + guardrail_to_apply=self, + buffer_until_moderated_default=False, + ) + ) as guarded: + async for streamed_chunk in guarded: + yield streamed_chunk return # Responses-API events are neither chat-completions chunks nor raw diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 37c1829def4..0c4ea6b29b5 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -6,11 +6,14 @@ Unified Guardrail, leveraging LiteLLM's /applyGuardrail endpoint 3. Implements a way to call /applyGuardrail endpoint for `/chat/completions` + `/v1/messages` requests on async_post_call_streaming_iterator_hook """ +import asyncio +import contextlib import copy import json from collections.abc import AsyncGenerator, AsyncIterable, Awaitable, Callable, Mapping, Sequence -from typing import TYPE_CHECKING, Any, Final, Protocol +from typing import TYPE_CHECKING, Any, Final, Protocol, TypeAlias, cast +import anyio from fastapi import HTTPException from litellm._logging import verbose_proxy_logger @@ -19,6 +22,7 @@ from litellm.cost_calculator import _infer_call_type from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route +from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket from litellm.llms import get_guardrail_translation_mapping, load_guardrail_translation_mappings from litellm.proxy._types import UserAPIKeyAuth from litellm.types.guardrails import GuardrailEventHooks @@ -28,6 +32,7 @@ from litellm.types.utils import ( CallTypesLiteral, Delta, ModelResponseStream, + StandardLoggingGuardrailInformation, StreamingChoices, ) @@ -45,6 +50,8 @@ A2A_CALL_TYPES: Final = (CallTypes.asend_message, CallTypes.send_message) GUARDRAIL_NAME: Final = "unified_llm_guardrails" +_RequestData: TypeAlias = dict[str, object] + class _EndpointTranslation(Protocol): @property @@ -59,6 +66,9 @@ class _EndpointTranslation(Protocol): @property def get_streaming_scan_key(self) -> "Callable[[Sequence[object]], StreamingScanKey | None]": ... + @property + def released_stream_as_ended(self) -> "Callable[[Sequence[object]], tuple[object, ...]]": ... + @property def build_block_sse_chunks(self) -> "Callable[..., Sequence[bytes] | None]": ... @@ -109,6 +119,18 @@ def _held_choices(held_chars_per_choice: Mapping[int, int]) -> frozenset[int]: return frozenset(idx for idx, held in held_chars_per_choice.items() if held > 0) +def _recorded_guardrail_information(request_data: _RequestData) -> tuple[StandardLoggingGuardrailInformation, ...]: + _metadata_key, metadata_bucket = get_or_create_metadata_bucket(request_data) + entries: Final = metadata_bucket.get("standard_logging_guardrail_information") + if not isinstance(entries, list): + return () + return tuple( + cast( # cast-ok: only the guardrail logging helpers write this metadata key + "list[StandardLoggingGuardrailInformation]", entries + ) + ) + + def _is_redundant_scan(scan_key: "StreamingScanKey | None", last_scan_key: "StreamingScanKey | None") -> bool: if scan_key is None: return False @@ -601,6 +623,7 @@ class UnifiedLLMGuardrails(CustomLogger): finish_reason_per_choice: dict[int, str | None], held_chars_per_choice: dict[int, int], is_final: bool, + terminated: asyncio.Event, ) -> AsyncGenerator[object, None]: """Run one guardrail processing round and emit the resulting diff chunk. @@ -632,6 +655,7 @@ class UnifiedLLMGuardrails(CustomLogger): is_final=is_final, ) except ModifyResponseException as e: + terminated.set() if e.original_response is None: e.original_response = responses_so_far async for block_chunk in self.handle_streaming_block( @@ -643,6 +667,7 @@ class UnifiedLLMGuardrails(CustomLogger): yield block_chunk raise _StreamTerminated() except HTTPException as e: + terminated.set() async for error_item in self.emit_streaming_http_error( e, call_type, @@ -664,7 +689,7 @@ class UnifiedLLMGuardrails(CustomLogger): *, guardrail_to_apply: CustomGuardrail, response: AsyncIterable[object], - request_data: dict, + request_data: _RequestData, user_api_key_dict: UserAPIKeyAuth, call_type: str, sampling_rate: int, @@ -687,6 +712,7 @@ class UnifiedLLMGuardrails(CustomLogger): held_chars_per_choice: Final[dict[int, int]] = {} chunk_counter = 0 last_chunk: object | None = None + terminated: Final = asyncio.Event() def _round(reference_chunk: object, is_final: bool) -> AsyncGenerator[object, None]: return self._emit_transform_round( @@ -702,10 +728,13 @@ class UnifiedLLMGuardrails(CustomLogger): finish_reason_per_choice=finish_reason_per_choice, held_chars_per_choice=held_chars_per_choice, is_final=is_final, + terminated=terminated, ) saw_tool_calls = False saw_text_content = False + tool_calls_released = False # rebind-ok: set once a raw tool call reaches the client unscanned + end_of_stream_inspection_started = False # rebind-ok: set once the end-of-stream inspection owns the verdict try: async for item in response: @@ -742,6 +771,7 @@ class UnifiedLLMGuardrails(CustomLogger): held_choices=_held_choices(held_chars_per_choice), ) responses_yielded.append(tool_only) + tool_calls_released = True yield tool_only continue @@ -781,6 +811,7 @@ class UnifiedLLMGuardrails(CustomLogger): # ``stream_transform_underflow`` 400 from mismatched prefixes. A shallow # list copy wouldn't help — the mutation is on the chunk objects # themselves — so we deepcopy. + end_of_stream_inspection_started = True if saw_tool_calls: async for out in self._inspect_full_response_for_block( endpoint_translation=endpoint_translation, @@ -801,6 +832,37 @@ class UnifiedLLMGuardrails(CustomLogger): yield out except _StreamTerminated: return + except (GeneratorExit, asyncio.CancelledError): + await self._scan_uninspected_tool_calls_after_disconnect( + uninspected=tool_calls_released and not end_of_stream_inspection_started and not terminated.is_set(), + endpoint_translation=endpoint_translation, + responses_released=responses_yielded, + guardrail_to_apply=guardrail_to_apply, + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) + raise + + @staticmethod + async def _scan_uninspected_tool_calls_after_disconnect( + *, + uninspected: bool, + endpoint_translation: _EndpointTranslation, + responses_released: Sequence[object], + guardrail_to_apply: CustomGuardrail, + user_api_key_dict: UserAPIKeyAuth, + request_data: _RequestData, + ) -> None: + if not uninspected: + return + await UnifiedLLMGuardrails._scan_released_stream_after_disconnect( + endpoint_translation=endpoint_translation, + responses_released=responses_released, + last_scan_key=None, + guardrail_to_apply=guardrail_to_apply, + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) async def _emit_stream_tail( self, @@ -841,14 +903,15 @@ class UnifiedLLMGuardrails(CustomLogger): from litellm.integrations.custom_guardrail import ModifyResponseException try: - await endpoint_translation.process_output_streaming_response( - responses_so_far=responses_so_far, - guardrail_to_apply=guardrail_to_apply, - litellm_logging_obj=request_data.get("litellm_logging_obj"), - user_api_key_dict=user_api_key_dict, - request_data=request_data, - stream_transform_sink=None, - ) + with anyio.CancelScope(shield=bool(responses_yielded)): + await endpoint_translation.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail_to_apply, + litellm_logging_obj=request_data.get("litellm_logging_obj"), + user_api_key_dict=user_api_key_dict, + request_data=request_data, + stream_transform_sink=None, + ) except ModifyResponseException as e: if e.original_response is None: e.original_response = responses_so_far @@ -968,6 +1031,45 @@ class UnifiedLLMGuardrails(CustomLogger): config_value: Final = config.get(name, attribute_value) if isinstance(config, dict) else attribute_value return self.optional_params.get(name, config_value) + @staticmethod + async def _scan_released_stream_after_disconnect( + *, + endpoint_translation: _EndpointTranslation, + responses_released: Sequence[object], + last_scan_key: "StreamingScanKey | None", + guardrail_to_apply: CustomGuardrail, + user_api_key_dict: UserAPIKeyAuth, + request_data: _RequestData, + ) -> None: + scanned: Final = endpoint_translation.released_stream_as_ended(copy.deepcopy(tuple(responses_released))) + if _is_redundant_scan(endpoint_translation.get_streaming_scan_key(scanned), last_scan_key): + return + recorded_before: Final = len(_recorded_guardrail_information(request_data)) + with anyio.CancelScope(shield=True): + try: + await endpoint_translation.process_output_streaming_response( + responses_so_far=scanned, + guardrail_to_apply=guardrail_to_apply, + litellm_logging_obj=request_data.get("litellm_logging_obj"), + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) + except Exception as e: # noqa: BLE001 # the client is gone, so the verdict can only be recorded + verbose_proxy_logger.warning( + "UnifiedLLMGuardrails: %s scanned a stream the client disconnected from and raised %s", + guardrail_to_apply.guardrail_name, + type(e).__name__, + ) + recorded_during_scan: Final = _recorded_guardrail_information(request_data)[recorded_before:] + if any(entry.get("guardrail_status") != "success" for entry in recorded_during_scan): + return + guardrail_to_apply.add_standard_logging_guardrail_information_to_request_data( + guardrail_json_response=e, + request_data=request_data, + guardrail_status="guardrail_failed_to_respond", + event_type=GuardrailEventHooks.post_call, + ) + async def async_post_call_streaming_iterator_hook( self, user_api_key_dict: UserAPIKeyAuth, @@ -995,6 +1097,7 @@ class UnifiedLLMGuardrails(CustomLogger): if guardrail_to_apply is None: guardrail_to_apply = request_data.pop("guardrail_to_apply", None) + typed_request_data: Final[_RequestData] = request_data def _streaming_flag(name: str, default: object) -> Any: return self.resolve_streaming_flag(guardrail_to_apply, name, default) @@ -1061,17 +1164,20 @@ class UnifiedLLMGuardrails(CustomLogger): mappings=mappings, ) if transform_call_type is not None: - async for transformed_item in self._run_incremental_transform_stream( - guardrail_to_apply=guardrail_to_apply, - response=response, - request_data=request_data, - user_api_key_dict=user_api_key_dict, - call_type=transform_call_type, - sampling_rate=sampling_rate, - end_of_stream_only=end_of_stream_only, - mappings=mappings, - ): - yield transformed_item + async with contextlib.aclosing( + self._run_incremental_transform_stream( + guardrail_to_apply=guardrail_to_apply, + response=response, + request_data=typed_request_data, + user_api_key_dict=user_api_key_dict, + call_type=transform_call_type, + sampling_rate=sampling_rate, + end_of_stream_only=end_of_stream_only, + mappings=mappings, + ) + ) as transformed: + async for transformed_item in transformed: + yield transformed_item return verbose_proxy_logger.warning( "UnifiedLLMGuardrails: streaming_transform_mode=incremental_diff is only supported " @@ -1093,217 +1199,240 @@ class UnifiedLLMGuardrails(CustomLogger): chunks_yielded = False last_scan_key: StreamingScanKey | None = None # rebind-ok: replaced after every scan round tool_calls_in_flight = False # rebind-ok: tracks the latest scan key's unscanned tool calls + verdict_settled = False # rebind-ok: set once the end-of-stream scan or a block owns the verdict - async for item in response: - chunk_counter += 1 - responses_so_far.append(item) + try: + async for item in response: + chunk_counter += 1 + responses_so_far.append(item) - # Infer call type from first chunk if not already done - if call_type is None and user_api_key_dict.request_route is not None: - call_types = get_call_types_for_route(user_api_key_dict.request_route) - if call_types is not None: - call_type = call_types[0].value + # Infer call type from first chunk if not already done + if call_type is None and user_api_key_dict.request_route is not None: + call_types = get_call_types_for_route(user_api_key_dict.request_route) + if call_types is not None: + call_type = call_types[0].value - if call_type is None: - call_type = _infer_call_type(call_type=None, completion_response=item) + if call_type is None: + call_type = _infer_call_type(call_type=None, completion_response=item) - # If call type not supported, just pass through all chunks - if call_type is None or CallTypes(call_type) not in mappings: - yield item - async for remaining_item in response: - yield remaining_item - return + # If call type not supported, just pass through all chunks + if call_type is None or CallTypes(call_type) not in mappings: + yield item + async for remaining_item in response: + yield remaining_item + return - # If end_of_stream_only mode, yield chunks without processing. - # When buffering, withhold them instead -- they are released (or - # replaced by the block message) only after end-of-stream - # moderation runs below. - if end_of_stream_only: - if not buffer_until_moderated: - endpoint_translation = mappings[CallTypes(call_type)]() - stream_has_ended = hasattr( - endpoint_translation, "_check_streaming_has_ended" - ) and endpoint_translation._check_streaming_has_ended(responses_so_far) - if pending_end_of_stream_items or stream_has_ended: - pending_end_of_stream_items.append(item) - else: - chunks_yielded = True - responses_yielded.append(item) - yield item - else: - withheld_items.append(item) - continue - - # Process chunk based on sampling rate - if buffer_until_moderated: - withheld_items.append(item) - if chunk_counter % sampling_rate == 0: - endpoint_translation = mappings[CallTypes(call_type)]() - scan_key = endpoint_translation.get_streaming_scan_key(responses_so_far) - if scan_key is not None: - tool_calls_in_flight = scan_key.tool_calls_in_flight - hold_window = buffer_until_moderated and (scan_key is None or tool_calls_in_flight) - if _is_redundant_scan(scan_key, last_scan_key): - verbose_proxy_logger.debug( - "Skipping streaming chunk %s for guardrail %s: nothing new to scan since the last round", - chunk_counter, - guardrail_to_apply.guardrail_name, - ) - if buffer_until_moderated: - if hold_window: - continue - for withheld_item in withheld_items: + # If end_of_stream_only mode, yield chunks without processing. + # When buffering, withhold them instead -- they are released (or + # replaced by the block message) only after end-of-stream + # moderation runs below. + if end_of_stream_only: + if not buffer_until_moderated: + endpoint_translation = mappings[CallTypes(call_type)]() + stream_has_ended = hasattr( + endpoint_translation, "_check_streaming_has_ended" + ) and endpoint_translation._check_streaming_has_ended(responses_so_far) + if pending_end_of_stream_items or stream_has_ended: + pending_end_of_stream_items.append(item) + else: chunks_yielded = True - responses_yielded.append(withheld_item) - yield withheld_item - withheld_items.clear() + responses_yielded.append(item) + yield item else: - chunks_yielded = True - responses_yielded.append(item) - yield item + withheld_items.append(item) continue + # Process chunk based on sampling rate + if buffer_until_moderated: + withheld_items.append(item) + if chunk_counter % sampling_rate == 0: + endpoint_translation = mappings[CallTypes(call_type)]() + scan_key = endpoint_translation.get_streaming_scan_key(responses_so_far) + if scan_key is not None: + tool_calls_in_flight = scan_key.tool_calls_in_flight + hold_window = buffer_until_moderated and (scan_key is None or tool_calls_in_flight) + if _is_redundant_scan(scan_key, last_scan_key): + verbose_proxy_logger.debug( + "Skipping streaming chunk %s for guardrail %s: nothing new to scan since the last round", + chunk_counter, + guardrail_to_apply.guardrail_name, + ) + if buffer_until_moderated: + if hold_window: + continue + for withheld_item in withheld_items: + chunks_yielded = True + responses_yielded.append(withheld_item) + yield withheld_item + withheld_items.clear() + else: + chunks_yielded = True + responses_yielded.append(item) + yield item + continue + + verbose_proxy_logger.debug( + "Processing streaming chunk %s (sampling_rate=%s) with guardrail %s", + chunk_counter, + sampling_rate, + guardrail_to_apply.guardrail_name, + ) + + original_items = ( + tuple(copy.deepcopy(withheld_items)) if buffer_until_moderated else (copy.deepcopy(item),) + ) + + try: + await endpoint_translation.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail_to_apply, + litellm_logging_obj=request_data.get("litellm_logging_obj"), + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) + except ModifyResponseException as e: + verdict_settled = True + if e.original_response is None: + e.original_response = responses_so_far + # Guardrail blocked the response mid-stream. Emit a clean + # terminating SSE sequence delivering the block message + # instead of letting the exception propagate into a bare + # `data: {"error": ...}` blob (which truncates the stream). + # Chunks have already been forwarded here, so the block + # continues the in-progress message (stream_started=True). + # The current chunk was appended to responses_so_far but not + # yet yielded, so exclude it: the continuation must reflect + # only what the client has actually received. + async for block_chunk in self.handle_streaming_block( + e, + endpoint_translation, + stream_started=chunks_yielded, + responses_so_far=responses_yielded, + ): + yield block_chunk + return + except HTTPException as e: + verdict_settled = True + # Response already started (we already yielded chunks); cannot send 400. + async for error_item in self.emit_streaming_http_error( + e, + call_type, + responses_so_far, + request_data, + endpoint_translation=endpoint_translation, + stream_started=chunks_yielded, + responses_yielded=responses_yielded, + ): + yield error_item + return + if scan_key is not None: + last_scan_key = scan_key + if hold_window: + verbose_proxy_logger.debug( + "Holding %s buffered chunks for guardrail %s: this round could not scan the whole window", + len(withheld_items), + guardrail_to_apply.guardrail_name, + ) + withheld_items[:] = original_items + continue + for original_item in original_items: + chunks_yielded = True + responses_yielded.append(original_item) + yield original_item + withheld_items.clear() + else: + if not buffer_until_moderated: + chunks_yielded = True + responses_yielded.append(item) + yield item + + # Stream has ended - do final processing with all collected chunks + if call_type is not None and CallTypes(call_type) in mappings: verbose_proxy_logger.debug( - "Processing streaming chunk %s (sampling_rate=%s) with guardrail %s", - chunk_counter, - sampling_rate, + "Processing final streaming response with all %s chunks for guardrail %s", + len(responses_so_far), guardrail_to_apply.guardrail_name, ) - original_items = ( - tuple(copy.deepcopy(withheld_items)) if buffer_until_moderated else (copy.deepcopy(item),) + endpoint_translation = mappings[CallTypes(call_type)]() + + buffered_items: Final = ( + tuple(copy.deepcopy(withheld_items)) + if buffer_until_moderated and release_on_scan and not end_of_stream_only + else tuple(withheld_items) + if buffer_until_moderated + else None ) + end_scan_key: Final = endpoint_translation.get_streaming_scan_key(responses_so_far) + verdict_settled = True + if _is_redundant_scan(end_scan_key, last_scan_key): + verbose_proxy_logger.debug( + "Skipping end-of-stream scan for guardrail %s: the last sampled round already scanned it all", + guardrail_to_apply.guardrail_name, + ) + for buffered_item in buffered_items or (): + yield buffered_item + for pending_item in pending_end_of_stream_items: + responses_yielded.append(pending_item) + yield pending_item + return try: - await endpoint_translation.process_output_streaming_response( - responses_so_far=responses_so_far, - guardrail_to_apply=guardrail_to_apply, - litellm_logging_obj=request_data.get("litellm_logging_obj"), - user_api_key_dict=user_api_key_dict, - request_data=request_data, - ) + with anyio.CancelScope(shield=chunks_yielded): + await endpoint_translation.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail_to_apply, + litellm_logging_obj=request_data.get("litellm_logging_obj"), + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) + # Moderation passed: release the withheld original chunks. + if buffered_items is not None: + for buffered_item in buffered_items: + yield buffered_item + for pending_item in pending_end_of_stream_items: + responses_yielded.append(pending_item) + yield pending_item except ModifyResponseException as e: if e.original_response is None: e.original_response = responses_so_far - # Guardrail blocked the response mid-stream. Emit a clean - # terminating SSE sequence delivering the block message - # instead of letting the exception propagate into a bare - # `data: {"error": ...}` blob (which truncates the stream). - # Chunks have already been forwarded here, so the block - # continues the in-progress message (stream_started=True). - # The current chunk was appended to responses_so_far but not - # yet yielded, so exclude it: the continuation must reflect - # only what the client has actually received. + # Block detected during end-of-stream processing. Emit a clean + # terminating SSE sequence with the block message rather than + # propagating into a bare error blob that truncates the stream. + # The withheld original chunks are never released. async for block_chunk in self.handle_streaming_block( e, endpoint_translation, - stream_started=chunks_yielded, + stream_started=bool(responses_yielded), responses_so_far=responses_yielded, ): yield block_chunk return except HTTPException as e: - # Response already started (we already yielded chunks); cannot send 400. async for error_item in self.emit_streaming_http_error( e, call_type, responses_so_far, request_data, endpoint_translation=endpoint_translation, - stream_started=chunks_yielded, + stream_started=bool(responses_yielded), responses_yielded=responses_yielded, ): yield error_item - return - if scan_key is not None: - last_scan_key = scan_key - if hold_window: - verbose_proxy_logger.debug( - "Holding %s buffered chunks for guardrail %s: this round could not scan the whole window", - len(withheld_items), - guardrail_to_apply.guardrail_name, - ) - withheld_items[:] = original_items - continue - for original_item in original_items: - chunks_yielded = True - responses_yielded.append(original_item) - yield original_item - withheld_items.clear() - else: - if not buffer_until_moderated: - chunks_yielded = True - responses_yielded.append(item) - yield item - - # Stream has ended - do final processing with all collected chunks - if call_type is not None and CallTypes(call_type) in mappings: - verbose_proxy_logger.debug( - "Processing final streaming response with all %s chunks for guardrail %s", - len(responses_so_far), - guardrail_to_apply.guardrail_name, - ) - - endpoint_translation = mappings[CallTypes(call_type)]() - - buffered_items: Final = ( - tuple(copy.deepcopy(withheld_items)) - if buffer_until_moderated and release_on_scan and not end_of_stream_only - else tuple(withheld_items) - if buffer_until_moderated - else None - ) - end_scan_key: Final = endpoint_translation.get_streaming_scan_key(responses_so_far) - if _is_redundant_scan(end_scan_key, last_scan_key): - verbose_proxy_logger.debug( - "Skipping end-of-stream scan for guardrail %s: the last sampled round already scanned it all", - guardrail_to_apply.guardrail_name, - ) - for buffered_item in buffered_items or (): - yield buffered_item - for pending_item in pending_end_of_stream_items: - responses_yielded.append(pending_item) - yield pending_item - return - - try: - await endpoint_translation.process_output_streaming_response( - responses_so_far=responses_so_far, + except (GeneratorExit, asyncio.CancelledError): + translation_class: Final = None if call_type is None else mappings.get(CallTypes(call_type)) + if ( + chunks_yielded + and not verdict_settled + and translation_class is not None + and isinstance(guardrail_to_apply, CustomGuardrail) + ): + await self._scan_released_stream_after_disconnect( + endpoint_translation=translation_class(), + responses_released=responses_yielded, + last_scan_key=last_scan_key, guardrail_to_apply=guardrail_to_apply, - litellm_logging_obj=request_data.get("litellm_logging_obj"), user_api_key_dict=user_api_key_dict, - request_data=request_data, + request_data=typed_request_data, ) - # Moderation passed: release the withheld original chunks. - if buffered_items is not None: - for buffered_item in buffered_items: - yield buffered_item - for pending_item in pending_end_of_stream_items: - responses_yielded.append(pending_item) - yield pending_item - except ModifyResponseException as e: - if e.original_response is None: - e.original_response = responses_so_far - # Block detected during end-of-stream processing. Emit a clean - # terminating SSE sequence with the block message rather than - # propagating into a bare error blob that truncates the stream. - # The withheld original chunks are never released. - async for block_chunk in self.handle_streaming_block( - e, - endpoint_translation, - stream_started=bool(responses_yielded), - responses_so_far=responses_yielded, - ): - yield block_chunk - return - except HTTPException as e: - async for error_item in self.emit_streaming_http_error( - e, - call_type, - responses_so_far, - request_data, - endpoint_translation=endpoint_translation, - stream_started=bool(responses_yielded), - responses_yielded=responses_yielded, - ): - yield error_item + raise diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index af487cb17a3..ed2a5b89a32 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -399,6 +399,7 @@ from litellm.proxy.common_request_processing import ( ProxyBaseLLMRequestProcessing, _is_azure_model_router_request, _should_return_raw_model_name, + close_guarded_stream, create_response, log_llm_api_exception, open_sse_before_first_byte, @@ -9911,6 +9912,17 @@ async def async_data_generator( stream_completed = False client_disconnected = False error_state: Final = ResponsesStreamErrorState() if responses_stream_errors else None + needs_iterator_wrap: Final = proxy_logging_obj.needs_iterator_wrap() + stream_iterator: Final[AsyncIterator[object]] = ( + proxy_logging_obj.async_post_call_streaming_iterator_hook( + user_api_key_dict=user_api_key_dict, + response=response, + request_data=request_data, + ) + if needs_iterator_wrap + else response + ) + stream_source: AsyncIterator[object] | None = None # rebind-ok: bound once the keepalive policy resolves try: error_message: str | None = None requested_model_from_client: Final = _get_client_requested_model_for_streaming(request_data=request_data) @@ -9935,21 +9947,11 @@ async def async_data_generator( # per-chunk hook. Coalescing them into a single flag forced wasted # ``get_response_string`` work per chunk on every deployment that # happened to ship a streaming-iterator override (the default). - needs_iterator_wrap: Final = proxy_logging_obj.needs_iterator_wrap() needs_per_chunk_hook: Final = proxy_logging_obj.needs_per_chunk_streaming_hook() is_raw_sse_stream: Final = bool(request_data.get("_litellm_raw_sse_stream")) strip_stream_usage: Final = bool(request_data.get("_litellm_strip_stream_usage")) raw_sse_buffer = "" - if needs_iterator_wrap: - stream_iterator = proxy_logging_obj.async_post_call_streaming_iterator_hook( - user_api_key_dict=user_api_key_dict, - response=response, - request_data=request_data, - ) - else: - stream_iterator = response - # A stream can start on a deployment with keepalive off and fall back # mid-stream to one that enables it: only skip wrapping altogether when # there's no router to ever fall back through AND the resolved interval @@ -9958,7 +9960,7 @@ async def async_data_generator( # happens to start with it off. resolve_keepalive_seconds: Final = _make_keepalive_resolver(request_data) initial_keepalive_seconds: Final = resolve_keepalive_seconds(response) - stream_source: Final = ( + stream_source = ( _iter_with_keepalive( stream_iterator.__aiter__(), resolve_keepalive_seconds, @@ -10099,6 +10101,9 @@ async def async_data_generator( # (a nested iterator hook would only see GeneratorExit on GC). if not stream_completed: client_disconnected = True + for guarded_layer in (stream_source, stream_iterator): + if guarded_layer is not response: + await close_guarded_stream(guarded_layer) raise except Exception as e: verbose_proxy_logger.exception("litellm.proxy.proxy_server.async_data_generator(): Exception occured - %s", e) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 1c738aefb97..29f2f46f001 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -44,6 +44,7 @@ from typing import ( Union, cast, overload, + runtime_checkable, ) from typing_extensions import ReadOnly, TypedDict @@ -524,8 +525,13 @@ class _UpstreamStreamBoundary(Generic[_T]): raise +@runtime_checkable +class _ClosableAsyncIterator(Protocol): + def aclose(self) -> object: ... + + class _StreamIteratorHook(Protocol[_T]): - def __call__(self, *, response: AsyncIterator[_T]) -> AsyncGenerator[_T, None]: ... + def __call__(self, *, response: AsyncIterator[_T]) -> AsyncIterator[_T]: ... def _is_client_error_exception(exc: Exception) -> bool: @@ -2773,8 +2779,22 @@ class ProxyLogging: ) -> AsyncGenerator[_T, None]: upstream: Final = _UpstreamStreamBoundary(response) try: - async for chunk in hook(response=upstream): - yield chunk + guarded: Final = hook(response=upstream) + try: + async for chunk in guarded: + yield chunk + finally: + if isinstance(guarded, _ClosableAsyncIterator): + try: + closing: Final = guarded.aclose() + if inspect.isawaitable(closing): + await closing + except Exception as e: # noqa: BLE001 # a finished stream must not fail on callback cleanup + verbose_proxy_logger.warning( + "Closing the streaming iterator of %s raised %s", + getattr(callback, "guardrail_name", None) or type(callback).__name__, + type(e).__name__, + ) except Exception as e: if e is not upstream.failure: enrich_http_exception_with_guardrail_context(e, callback) @@ -3959,6 +3979,7 @@ class ProxyLogging: stream_needs_translation: Final = ProxyLogging._stream_requires_guardrail_translation(user_api_key_dict) pipeline_gated_names: Final = _pipeline_step_guardrail_names(post_call_pipelines) + guarded_layers: Final[list[AsyncGenerator[object, None]]] = [] # mutable-ok: closed on disconnect for resolved_callback, kind in caps.iterator_overrides: if isinstance(resolved_callback, CustomGuardrail): if resolved_callback.guardrail_name in pipeline_gated_names: @@ -4001,6 +4022,7 @@ class ProxyLogging: hook, request_data=request_data, ) + guarded_layers.append(current_response) pipeline_translation: Final = ( resolve_endpoint_translation(user_api_key_dict, None) if post_call_pipelines else None @@ -4013,6 +4035,7 @@ class ProxyLogging: pipelines=post_call_pipelines, translation=pipeline_translation, ) + guarded_layers.append(current_response) served_chunks: Final[list[object]] = [] # mutable-ok: accumulates while yielding to the client try: @@ -4020,6 +4043,7 @@ class ProxyLogging: served_chunks.append(chunk) yield chunk except (GeneratorExit, asyncio.CancelledError): + await ProxyLogging._close_guarded_layers(guarded_layers) ProxyLogging._record_served_stream_output(request_data, served_chunks) raise except Exception as e: @@ -4100,6 +4124,16 @@ class ProxyLogging: for buffered_item in buffered: yield buffered_item + @staticmethod + async def _close_guarded_layers(layers: Sequence[AsyncGenerator[object, None]]) -> None: + for layer in reversed(layers): + try: + await layer.aclose() + except Exception as e: # noqa: BLE001 # one failing callback cleanup must not skip the inner ones + verbose_proxy_logger.warning( + "Closing a streaming callback layer after a client disconnect raised %s", type(e).__name__ + ) + @staticmethod def _record_served_stream_output(request_data: Mapping[str, object], served_chunks: Sequence[object]) -> None: logging_obj: Final = request_data.get("litellm_logging_obj") diff --git a/tests/integration/observability/test_guardrail_effects.py b/tests/integration/observability/test_guardrail_effects.py index d377afb206c..1df4427bb4e 100644 --- a/tests/integration/observability/test_guardrail_effects.py +++ b/tests/integration/observability/test_guardrail_effects.py @@ -1,12 +1,20 @@ +import asyncio import json import os +import re import signal import socket +import threading import uuid +from collections.abc import Callable, Iterator, Mapping from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager +from dataclasses import dataclass from pathlib import Path +from types import MappingProxyType from typing import Final +import anthropic import httpx import psutil import pytest @@ -14,9 +22,10 @@ import yaml from integration._support.client import Gateway, eventually, object_value from integration._support.database import read_rows from integration._support.mcp import mcp_peer, register_mcp, tool_names -from integration._support.process import group_members, owned_proxy, owned_proxy_process -from integration._support.wire import Reply, Request, wire_server +from integration._support.process import OwnedProxy, group_members, owned_proxy, owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server from openai import AsyncOpenAI, OpenAI +from pydantic import JsonValue @pytest.mark.covers("other.observability.guardrails.rewrite_reaches_correct_anthropic_positions") @@ -1796,3 +1805,739 @@ def test_responses_pre_call_denial_stream_survives_worker_kill(gateway: Gateway, for response in responses: assert response.status_code == 200, response.text assert response.headers["content-type"].startswith("text/event-stream"), response.text + + +_TOKEN: Final = re.compile(rb"token-[0-9a-f]{32}-\d+") + + +def _secret_for(request: Request) -> str: + token: Final = _TOKEN.search(request.body) + assert token is not None, request.body + return "synthetic-leaked-secret-" + token.group().decode() + + +def _chat_frame(identity: str, choices: tuple[dict[str, JsonValue], ...]) -> bytes: + payload: Final = { + "id": identity, + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o-mini", + "choices": list(choices), + } + return b"data: " + json.dumps(payload).encode() + b"\n\n" + + +def _chat_choice(index: int, delta: dict[str, JsonValue], finish: str | None = None) -> dict[str, JsonValue]: + return {"index": index, "delta": delta, "finish_reason": finish} + + +def _chat_stream_frames(secret: str, shape: str) -> tuple[bytes, ...]: + identity: Final = "chatcmpl-" + uuid.uuid4().hex + tool_call: Final = { + "index": 0, + "id": "call_" + identity, + "type": "function", + "function": {"name": "lookup", "arguments": json.dumps({"query": secret})}, + } + released: Final = { + "text": (_chat_choice(0, {"role": "assistant", "content": secret}),), + "empty": (_chat_choice(0, {"role": "assistant", "content": ""}),), + "tool_call": (_chat_choice(0, {"role": "assistant", "tool_calls": [tool_call]}),), + "two_choices": ( + _chat_choice(0, {"role": "assistant", "content": secret + "-first"}), + _chat_choice(1, {"role": "assistant", "content": secret + "-second"}), + ), + }[shape] + finish: Final = "tool_calls" if shape == "tool_call" else "stop" + tail: Final = tuple(_chat_choice(int(str(choice["index"])), {"content": " tail"}, finish) for choice in released) + return (_chat_frame(identity, released), _chat_frame(identity, tail), b"data: [DONE]\n\n") + + +def _chat_completion_body(secret: str) -> bytes: + return json.dumps( + { + "id": "chatcmpl-" + uuid.uuid4().hex, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": secret}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 9, "completion_tokens": 4, "total_tokens": 13}, + } + ).encode() + + +def _responses_tool_call_frames(secret: str, *, with_text: bool) -> tuple[bytes, ...]: + identity: Final = "resp_" + uuid.uuid4().hex + arguments: Final = json.dumps({"query": secret}) + pending: Final = {"type": "function_call", "id": "fc_" + identity, "call_id": "call_" + identity, "name": "lookup"} + finished: Final = {**pending, "arguments": arguments, "status": "completed"} + envelope: Final = {"id": identity, "object": "response", "created_at": 1, "model": "gpt-4o-mini", "output": []} + message: Final = { + "type": "message", + "id": "msg_" + identity, + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": secret, "annotations": []}], + } + text_delta: Final = { + "type": "response.output_text.delta", + "item_id": "msg_" + identity, + "output_index": 0, + "content_index": 0, + "delta": secret, + } + text_events: Final = (text_delta,) if with_text else () + tool_index: Final = len(text_events) + output: Final = [*((message,) if with_text else ()), finished] + events: Final = ( + {"type": "response.created", "response": {**envelope, "status": "in_progress"}}, + *text_events, + { + "type": "response.output_item.added", + "output_index": tool_index, + "item": {**pending, "arguments": "", "status": "in_progress"}, + }, + { + "type": "response.function_call_arguments.delta", + "item_id": "fc_" + identity, + "output_index": tool_index, + "delta": arguments, + }, + {"type": "response.output_item.done", "output_index": tool_index, "item": finished}, + {"type": "response.completed", "response": {**envelope, "status": "completed", "output": output}}, + ) + encoded: Final = tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events) + released_through: Final = len(events) - 1 if with_text else 3 + return (b"".join(encoded[:released_through]), b"".join(encoded[released_through:])) + + +def _responses_stream_frames(secret: str) -> tuple[bytes, ...]: + identity: Final = "resp_" + uuid.uuid4().hex + completed: Final = { + "id": identity, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "type": "message", + "id": "msg_" + identity, + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": secret, "annotations": []}], + } + ], + "usage": { + "input_tokens": 11, + "output_tokens": 4, + "total_tokens": 15, + "input_tokens_details": {"cached_tokens": 0}, + "output_tokens_details": {"reasoning_tokens": 0}, + }, + } + events: Final = ( + {"type": "response.created", "response": {**completed, "status": "in_progress", "output": [], "usage": None}}, + { + "type": "response.output_text.delta", + "item_id": "msg_" + identity, + "output_index": 0, + "content_index": 0, + "delta": secret, + }, + {"type": "response.completed", "response": completed}, + ) + encoded: Final = tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events) + return (encoded[0] + encoded[1], encoded[2]) + + +def _gemini_stream_frames(secret: str) -> tuple[bytes, ...]: + def frame(text: str, finish: str | None) -> bytes: + candidate: Final = { + "content": {"parts": [{"text": text}], "role": "model"}, + "index": 0, + **({"finishReason": finish} if finish else {}), + } + payload: Final = { + "candidates": [candidate], + "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 5, "totalTokenCount": 15}, + "modelVersion": "gemini-2.5-flash", + } + return b"data: " + json.dumps(payload).encode() + b"\r\n\r\n" + + return (frame(secret, None), frame(" tail", "STOP")) + + +def _scripted_provider(gate: threading.Event | None, pause: float, shape: str) -> Callable[[Request], Reply]: + def provider(request: Request) -> Reply: + secret: Final = _secret_for(request) + path: Final = request.target.split("?")[0] + if path.endswith("/chat/completions") and not json.loads(request.body).get("stream"): + return Reply(body=_chat_completion_body(secret)) + frames: Final = ( + _gemini_stream_frames(secret) + if "streamGenerateContent" in path + else ( + _responses_tool_call_frames(secret, with_text=shape == "text_then_tool_call") + if shape in ("tool_call", "text_then_tool_call") + else _responses_stream_frames(secret) + ) + if path.endswith("/responses") + else _chat_stream_frames(secret, shape) + ) + return Reply(content_type="text/event-stream", chunks=frames, gate_after_first=gate, pause_between_chunks=pause) + + return provider + + +def _allowing_guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=json.dumps({"action": "NONE"}).encode()) + + +def _failing_response_scans(reply: Reply) -> Callable[[Request], Reply]: + def guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + if json.loads(request.body)["input_type"] == "response": + return reply + return Reply(body=json.dumps({"action": "NONE"}).encode()) + + return guardrail + + +def _post_call_config( + tmp_path: Path, identity: str, policy_url: str, params: Mapping[str, JsonValue], default_on: bool +) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "post_call", + "default_on": default_on, + "api_base": policy_url, + "api_key": "synthetic-guardrail-key", + **params, + }, + } + ] + path: Final = tmp_path / f"{identity}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +class _ScanLog: + def __init__(self, policy: Wire) -> None: + self.policy: Final = policy + self.seen: tuple[dict[str, JsonValue], ...] = () + + def response_scans(self, secret: str) -> tuple[dict[str, JsonValue], ...]: + self.seen = (*self.seen, *(object_value(json.loads(request.body)) for request in self.policy.drain())) + return tuple(body for body in self.seen if body["input_type"] == "response" and secret in json.dumps(body)) + + +@dataclass(frozen=True, slots=True) +class _DisconnectRig: + owned: OwnedProxy + model: str + gemini: str + scans: _ScanLog + identity: str + gate: threading.Event + upstream: Wire + + @property + def candidate(self) -> Gateway: + return self.owned.gateway + + def token(self, index: int = 0) -> str: + return f"token-{self.identity.removeprefix('guardrail')}-{index}" + + def secret(self, index: int = 0) -> str: + return "synthetic-leaked-secret-" + self.token(index) + + +_END_OF_STREAM_ONLY: Final = MappingProxyType({"streaming_end_of_stream_only": True}) + + +@contextmanager +def _disconnect_rig( + gateway: Gateway, + tmp_path: Path, + *, + params: Mapping[str, JsonValue] = _END_OF_STREAM_ONLY, + guardrail: Callable[[Request], Reply] = _allowing_guardrail, + gated: bool = True, + pause: float = 0, + shape: str = "text", + default_on: bool = True, + workers: int = 1, +) -> Iterator[_DisconnectRig]: + identity: Final = "guardrail" + uuid.uuid4().hex + gate: Final = threading.Event() + with ( + wire_server(guardrail) as policy, + wire_server(_scripted_provider(gate if gated else None, pause, shape)) as upstream, + ): + config: Final = _post_call_config(tmp_path, identity, policy.url, params, default_on) + try: + with ( + owned_proxy_process(gateway, tmp_path, {}, config=config, workers=workers) as owned, + owned.gateway.scenario() as scenario, + ): + yield _DisconnectRig( + owned, + scenario.model(api_base=upstream.url + "/v1", api_key="synthetic-openai-key"), + scenario.model( + model="gemini/gemini-2.5-flash", api_base=upstream.url, api_key="synthetic-gemini-key" + ), + _ScanLog(policy), + identity, + gate, + upstream, + ) + finally: + gate.set() + + +def _close_on(response: httpx.Response, marker: str) -> str: + assert response.status_code == 200, response.read() + for line in response.iter_lines(): + if marker in line: + return line + raise AssertionError(f"The stream ended before the client received {marker}") + + +def _stream_and_close(rig: _DisconnectRig, path: str, body: Mapping[str, JsonValue], marker: str) -> str: + with rig.candidate.client.stream( + "POST", path, json=dict(body), headers={"Authorization": f"Bearer {rig.candidate.key}"} + ) as response: + return _close_on(response, marker) + + +def _chat_body(rig: _DisconnectRig, index: int, stream: bool = True) -> dict[str, JsonValue]: + return { + "model": rig.model, + "messages": [{"role": "user", "content": "synthetic prompt " + rig.token(index)}], + "stream": stream, + } + + +def _chat_httpx(rig: _DisconnectRig, index: int = 0) -> str: + return _stream_and_close(rig, "/v1/chat/completions", _chat_body(rig, index), rig.secret(index)) + + +def _responses_httpx(rig: _DisconnectRig, index: int = 0) -> str: + body: Final = {"model": rig.model, "input": "synthetic prompt " + rig.token(index), "stream": True} + return _stream_and_close(rig, "/v1/responses", body, rig.secret(index)) + + +def _messages_httpx(rig: _DisconnectRig, index: int = 0) -> str: + body: Final = {**_chat_body(rig, index), "max_tokens": 64} + return _stream_and_close(rig, "/v1/messages", body, rig.secret(index)) + + +def _gemini_httpx(rig: _DisconnectRig, index: int = 0) -> str: + body: Final = {"contents": [{"role": "user", "parts": [{"text": "synthetic prompt " + rig.token(index)}]}]} + path: Final = f"/v1beta/models/{rig.gemini}:streamGenerateContent?alt=sse" + return _stream_and_close(rig, path, body, rig.secret(index)) + + +def _chat_async_openai_sdk(rig: _DisconnectRig, index: int = 0) -> str: + async def read() -> str: + client: Final = AsyncOpenAI( + base_url=str(rig.candidate.client.base_url) + "/v1", + api_key=rig.candidate.key, + max_retries=0, + http_client=httpx.AsyncClient(trust_env=False, timeout=30), + ) + async with client: + stream: Final = await client.chat.completions.create( + model=rig.model, + messages=[{"role": "user", "content": "synthetic prompt " + rig.token(index)}], + stream=True, + ) + async for chunk in stream: + if chunk.choices and rig.secret(index) in (chunk.choices[0].delta.content or ""): + await stream.close() + return chunk.choices[0].delta.content or "" + raise AssertionError("The stream ended before the client received the streamed content") + + return asyncio.run(read()) + + +def _responses_openai_sdk(rig: _DisconnectRig, index: int = 0) -> str: + client: Final = OpenAI( + base_url=str(rig.candidate.client.base_url) + "/v1", + api_key=rig.candidate.key, + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=30), + ) + with client: + stream: Final = client.responses.create( + model=rig.model, input="synthetic prompt " + rig.token(index), stream=True + ) + for event in stream: + if event.type == "response.output_text.delta" and rig.secret(index) in event.delta: + stream.close() + return event.delta + raise AssertionError("The stream ended before the client received the streamed content") + + +def _messages_anthropic_sdk(rig: _DisconnectRig, index: int = 0) -> str: + client: Final = anthropic.Anthropic( + base_url=str(rig.candidate.client.base_url), + api_key=rig.candidate.key, + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=30), + ) + with client: + stream: Final = client.messages.create( + model=rig.model, + max_tokens=64, + messages=[{"role": "user", "content": "synthetic prompt " + rig.token(index)}], + stream=True, + ) + for event in stream: + text: Final = ( + event.delta.text if event.type == "content_block_delta" and event.delta.type == "text_delta" else "" + ) + if rig.secret(index) in text: + stream.close() + return text + raise AssertionError("The stream ended before the client received the streamed content") + + +def _scanned_while_upstream_is_held( + rig: _DisconnectRig, disconnect: Callable[[_DisconnectRig, int], str], index: int = 0 +) -> tuple[dict[str, JsonValue], ...]: + try: + received: Final = disconnect(rig, index) + assert rig.secret(index) in received, received + return eventually( + lambda: rig.scans.response_scans(rig.secret(index)), lambda values: len(values) >= 1, seconds=4 + ) + finally: + rig.gate.set() + + +def _post_call_statuses(rig: _DisconnectRig, model: str, rows: int = 1) -> tuple[tuple[str, ...], ...]: + found: Final = eventually( + lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)), + lambda values: len(values) == rows, + seconds=70, + ) + + def statuses(metadata: JsonValue) -> tuple[str, ...]: + entries: Final = object_value(metadata).get("guardrail_information") or [] + assert isinstance(entries, list), metadata + post_call: Final = tuple( + entry + for entry in (object_value(value) for value in entries) + if entry.get("guardrail_name") == rig.identity and entry.get("guardrail_mode") == "post_call" + ) + return tuple(str(entry["guardrail_status"]) for entry in post_call) + + return tuple(statuses(row["metadata"]) for row in found) + + +_ENDPOINT_CLIENTS: Final = ( + pytest.param(_chat_httpx, id="chat-httpx"), + pytest.param(_chat_async_openai_sdk, id="chat-async-openai-sdk"), + pytest.param(_responses_httpx, id="responses-httpx"), + pytest.param(_responses_openai_sdk, id="responses-openai-sdk"), + pytest.param(_messages_httpx, id="messages-httpx"), + pytest.param(_messages_anthropic_sdk, id="messages-anthropic-sdk"), + pytest.param(_gemini_httpx, id="native-gemini-stream-generate-content"), +) + + +_NO_DISCONNECT_ROW_LIT_8603: Final = pytest.mark.skip( + reason="BUG: LIT-8603 a mid-stream disconnect writes no spend row" +) + + +@pytest.mark.parametrize("disconnect", _ENDPOINT_CLIENTS) +def test_client_disconnect_mid_stream_still_scans_the_content_it_already_received( + gateway: Gateway, tmp_path: Path, disconnect: Callable[[_DisconnectRig, int], str] +) -> None: + with _disconnect_rig(gateway, tmp_path) as rig: + scans: Final = _scanned_while_upstream_is_held(rig, disconnect) + assert len(scans) == 1, scans + + +@pytest.mark.parametrize( + "disconnect", + ( + pytest.param(_chat_httpx, id="chat-httpx"), + pytest.param(_chat_async_openai_sdk, id="chat-async-openai-sdk"), + pytest.param(_responses_httpx, id="responses-httpx", marks=_NO_DISCONNECT_ROW_LIT_8603), + pytest.param(_messages_anthropic_sdk, id="messages-anthropic-sdk", marks=_NO_DISCONNECT_ROW_LIT_8603), + pytest.param( + _gemini_httpx, + id="native-gemini-stream-generate-content", + marks=pytest.mark.skip(reason="BUG: LIT-9087 a mid-stream disconnect writes no spend row"), + ), + ), +) +def test_client_disconnect_mid_stream_records_the_post_call_verdict_on_the_spend_row( + gateway: Gateway, tmp_path: Path, disconnect: Callable[[_DisconnectRig, int], str] +) -> None: + with _disconnect_rig(gateway, tmp_path) as rig: + _scanned_while_upstream_is_held(rig, disconnect) + model: Final = rig.gemini if disconnect is _gemini_httpx else rig.model + assert _post_call_statuses(rig, model) == (("success",),) + + +@pytest.mark.parametrize( + "reply", + ( + pytest.param(Reply(status=500, body=b'{"error": "synthetic guardrail outage"}'), id="guardrail-500"), + pytest.param(Reply(body=b"synthetic non-json guardrail body"), id="guardrail-malformed-200"), + ), +) +def test_client_disconnect_mid_stream_records_a_failed_scan_when_the_guardrail_errors( + gateway: Gateway, tmp_path: Path, reply: Reply +) -> None: + with _disconnect_rig(gateway, tmp_path, guardrail=_failing_response_scans(reply)) as rig: + _scanned_while_upstream_is_held(rig, _chat_httpx) + assert _post_call_statuses(rig, rig.model) == (("guardrail_failed_to_respond",),) + + +def test_client_disconnect_mid_stream_records_a_blocking_verdict_and_keeps_serving( + gateway: Gateway, tmp_path: Path +) -> None: + blocked: Final = Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic leak"}).encode()) + with _disconnect_rig(gateway, tmp_path, guardrail=_failing_response_scans(blocked)) as rig: + _scanned_while_upstream_is_held(rig, _chat_httpx) + statuses: Final = _post_call_statuses(rig, rig.model) + assert len(statuses) == 1 and len(statuses[0]) == 1 and statuses[0][0] != "success", statuses + health: Final = rig.candidate.request("GET", "/health/liveliness") + assert health.status_code == 200, health.text + + +def test_client_disconnect_while_end_of_stream_scan_is_in_flight_still_records_the_verdict( + gateway: Gateway, tmp_path: Path +) -> None: + scan_started: Final = threading.Event() + scan_released: Final = threading.Event() + + def guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + if json.loads(request.body)["input_type"] == "response": + scan_started.set() + assert scan_released.wait(timeout=30), "The in-flight scan was never released" + return Reply(body=json.dumps({"action": "NONE"}).encode()) + + with _disconnect_rig(gateway, tmp_path, guardrail=guardrail, gated=False) as rig: + with rig.candidate.client.stream( + "POST", + "/v1/chat/completions", + json=_chat_body(rig, 0), + headers={"Authorization": f"Bearer {rig.candidate.key}"}, + ) as response: + try: + assert rig.secret() in _close_on(response, rig.secret()) + assert scan_started.wait(timeout=10), "The end-of-stream scan never started" + finally: + pass + scan_released.set() + assert _post_call_statuses(rig, rig.model) == (("success",),) + assert len(rig.scans.response_scans(rig.secret())) == 1 + + +@pytest.mark.parametrize( + ("params", "shape", "expected"), + ( + pytest.param( + {"streaming_buffer_until_moderated": False}, + "text", + ("synthetic-leaked-secret-",), + id="sampled-before-the-sampling-threshold", + ), + pytest.param( + {"streaming_buffer_until_moderated": False, "streaming_transform_mode": "incremental_diff"}, + "tool_call", + ('\\"query\\": \\"synthetic-leaked-secret-',), + id="incremental-diff-tool-call-in-flight", + ), + pytest.param(dict(_END_OF_STREAM_ONLY), "two_choices", ("-first", "-second"), id="two-choices"), + ), +) +def test_client_disconnect_mid_stream_scans_what_each_streaming_mode_released( + gateway: Gateway, tmp_path: Path, params: dict[str, JsonValue], shape: str, expected: tuple[str, ...] +) -> None: + with _disconnect_rig(gateway, tmp_path, params=params, shape=shape) as rig: + marker: Final = rig.secret() + ("-first" if shape == "two_choices" else "") + try: + received: Final = _stream_and_close(rig, "/v1/chat/completions", _chat_body(rig, 0), rig.token()) + assert rig.token() in received, received + scans: Final = eventually( + lambda: rig.scans.response_scans(rig.secret()), lambda values: len(values) >= 1, seconds=4 + ) + finally: + rig.gate.set() + payload: Final = json.dumps(scans[-1]) + assert all(fragment in payload for fragment in expected), (marker, scans) + assert _post_call_statuses(rig, rig.model)[0][-1:] == ("success",) + + +def _tool_call_request(rig: _DisconnectRig, path: str) -> dict[str, JsonValue]: + if path == "/v1/responses": + return {"model": rig.model, "input": "synthetic prompt " + rig.token(), "stream": True} + if path == "/v1/messages": + return {**_chat_body(rig, 0), "max_tokens": 64} + return _chat_body(rig, 0) + + +@pytest.mark.parametrize( + "path", + ( + pytest.param("/v1/chat/completions", id="chat"), + pytest.param("/v1/responses", id="responses"), + pytest.param("/v1/messages", id="messages"), + ), +) +def test_client_disconnect_mid_tool_call_scans_the_tool_call_it_already_received( + gateway: Gateway, tmp_path: Path, path: str +) -> None: + with _disconnect_rig(gateway, tmp_path, shape="tool_call") as rig: + try: + received: Final = _stream_and_close(rig, path, _tool_call_request(rig, path), rig.token()) + assert rig.token() in received, received + scans: Final = eventually( + lambda: rig.scans.response_scans(rig.secret()), lambda values: len(values) >= 1, seconds=4 + ) + finally: + rig.gate.set() + assert rig.secret() in json.dumps(scans[-1].get("tool_calls")), scans + + +def test_client_disconnect_after_a_finished_responses_tool_call_scans_the_text_and_tool_call_it_received( + gateway: Gateway, tmp_path: Path +) -> None: + with _disconnect_rig(gateway, tmp_path, shape="text_then_tool_call") as rig: + try: + received: Final = _stream_and_close( + rig, "/v1/responses", _tool_call_request(rig, "/v1/responses"), "response.output_item.done" + ) + assert rig.token() in received, received + scans: Final = eventually( + lambda: tuple(scan for scan in rig.scans.response_scans(rig.secret()) if scan.get("texts")), + lambda values: len(values) >= 1, + seconds=4, + ) + finally: + rig.gate.set() + assert rig.secret() in json.dumps(scans[-1].get("texts")), scans + assert rig.secret() in json.dumps(scans[-1].get("tool_calls")), scans + + +def test_client_disconnect_mid_stream_scans_for_a_guardrail_the_request_opted_into( + gateway: Gateway, tmp_path: Path +) -> None: + with _disconnect_rig(gateway, tmp_path, default_on=False) as rig: + body: Final = {**_chat_body(rig, 0), "guardrails": [rig.identity]} + try: + received: Final = _stream_and_close(rig, "/v1/chat/completions", body, rig.secret()) + assert rig.secret() in received, received + eventually(lambda: rig.scans.response_scans(rig.secret()), lambda values: len(values) == 1, seconds=4) + finally: + rig.gate.set() + assert _post_call_statuses(rig, rig.model) == (("success",),) + + +def test_client_disconnect_before_any_content_sends_no_response_scan(gateway: Gateway, tmp_path: Path) -> None: + with _disconnect_rig(gateway, tmp_path, shape="empty") as rig: + try: + _stream_and_close(rig, "/v1/chat/completions", _chat_body(rig, 0), "data: ") + finally: + rig.gate.set() + rows: Final = _post_call_statuses(rig, rig.model) + assert rig.scans.response_scans(rig.secret()) == (), rig.scans.seen + assert len(rows) == 1 and "success" not in rows[0], rows + + +@pytest.mark.parametrize( + "params", + ( + pytest.param(dict(_END_OF_STREAM_ONLY), id="end-of-stream-only"), + pytest.param({"streaming_buffer_until_moderated": True}, id="buffered"), + ), +) +def test_a_fully_read_stream_is_scanned_exactly_once( + gateway: Gateway, tmp_path: Path, params: dict[str, JsonValue] +) -> None: + with _disconnect_rig(gateway, tmp_path, params=params, gated=False) as rig: + response: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig, 0)) + assert response.status_code == 200, response.text + assert rig.secret() in response.text and "[DONE]" in response.text, response.text + assert _post_call_statuses(rig, rig.model) == (("success",),) + assert len(rig.scans.response_scans(rig.secret())) == 1, rig.scans.seen + + +def _cached_twin_rows(rig: _DisconnectRig) -> tuple[tuple[str, ...], ...]: + first: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig, 0, stream=False)) + second: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig, 0, stream=False)) + assert first.status_code == second.status_code == 200, (first.text, second.text) + assert first.json()["choices"] == second.json()["choices"], (first.text, second.text) + assert len(rig.upstream.drain()) == 1, "the second request must be served from the cache" + return _post_call_statuses(rig, rig.model, rows=2) + + +def test_a_non_streaming_response_and_its_cache_hit_are_each_scanned_once(gateway: Gateway, tmp_path: Path) -> None: + with _disconnect_rig(gateway, tmp_path, gated=False) as rig: + rows: Final = _cached_twin_rows(rig) + assert rows[0] == ("success",), rows + assert len(rig.scans.response_scans(rig.secret())) == 2, (rows, rig.scans.seen) + + +def test_a_cache_hit_row_records_the_post_call_verdict_of_its_scan(gateway: Gateway, tmp_path: Path) -> None: + pytest.skip("BUG: LIT-9088 the cache-hit spend row drops the post_call verdict of the scan that ran on it") + with _disconnect_rig(gateway, tmp_path, gated=False) as rig: + assert _cached_twin_rows(rig) == (("success",), ("success",)) + + +def test_concurrent_disconnects_during_a_guardrail_outage_each_record_exactly_one_verdict( + gateway: Gateway, tmp_path: Path +) -> None: + def guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + body: Final = json.loads(request.body) + index: Final = int(_secret_for(request).rsplit("-", 1)[1]) + if body["input_type"] == "response" and index % 3 == 0: + return Reply(status=503, body=b'{"error": "synthetic guardrail outage"}') + return Reply(body=json.dumps({"action": "NONE"}).encode()) + + clients: Final = (_chat_httpx, _responses_httpx, _messages_httpx) + with _disconnect_rig(gateway, tmp_path, guardrail=guardrail, gated=False, pause=3, workers=2) as rig: + with ThreadPoolExecutor(max_workers=30) as pool: + received: Final = tuple(pool.map(lambda index: clients[index % 3](rig, index), range(30))) + assert all(rig.secret(index) in line for index, line in enumerate(received)), received + scanned: Final = eventually( + lambda: tuple(len(rig.scans.response_scans(rig.secret(index) + '"')) for index in range(30)), + lambda counts: all(count >= 1 for count in counts), + seconds=20, + ) + assert scanned == (1,) * 30, scanned + chat_rows: Final = _post_call_statuses(rig, rig.model, rows=10) + assert sorted(chat_rows) == sorted( + ("guardrail_failed_to_respond",) if index % 3 == 0 else ("success",) for index in range(0, 30, 3) + ), chat_rows + + +def test_disconnect_scans_keep_recording_after_a_worker_is_killed(gateway: Gateway, tmp_path: Path) -> None: + with _disconnect_rig(gateway, tmp_path, gated=False, pause=3, workers=2) as rig: + members: Final = tuple( + member for member in group_members(rig.owned.process.pid) if member.pid != rig.owned.process.pid + ) + workers: Final = tuple(member for member in members if any("spawn_main" in part for part in member.cmdline())) + assert len(workers) >= 2, members + workers[0].send_signal(signal.SIGKILL) + psutil.wait_procs((workers[0],), timeout=10) + with ThreadPoolExecutor(max_workers=8) as pool: + received: Final = tuple(pool.map(lambda index: _chat_httpx(rig, index), range(8))) + assert all(rig.secret(index) in line for index, line in enumerate(received)), received + assert _post_call_statuses(rig, rig.model, rows=8) == (("success",),) * 8 + assert rig.owned.process.poll() is None diff --git a/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py index 09921a9a85c..85e4c055f07 100644 --- a/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py +++ b/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py @@ -2489,6 +2489,28 @@ class TestAnthropicMessagesHandlerStreamingScanKey: assert ended_key.tool_calls_in_flight is False assert ended_key != open_key + def test_released_stream_as_ended_keys_the_tool_use_the_client_already_received(self): + handler = AnthropicMessagesHandler() + tool_use = self._sse( + "content_block_start", + { + "type": "content_block_start", + "index": 1, + "content_block": {"type": "tool_use", "id": "toolu_1", "name": "get_weather", "input": {}}, + }, + ) + stopped_key = handler.get_streaming_scan_key([self._text_delta("hi"), tool_use, self._stop("tool_use")]) + released_key = handler.get_streaming_scan_key( + handler.released_stream_as_ended([self._text_delta("hi"), tool_use]) + ) + assert released_key.stream_ended is True + assert released_key == stopped_key + + def test_released_stream_as_ended_leaves_a_text_only_stream_as_released(self): + released = (self._text_delta("hi"), self._text_delta(" there")) + ended = AnthropicMessagesHandler().released_stream_as_ended(released) + assert ended == released and all(a is b for a, b in zip(ended, released, strict=True)) + class PerRowTextGuardrail(CustomGuardrail): """Answers one redacted text per chat row it was shown, the way a guardrail diff --git a/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py b/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py index 5c85faa5e13..5d71cbe9afc 100644 --- a/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py +++ b/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py @@ -2336,3 +2336,32 @@ class TestStreamingScanKey: handler = OpenAIChatCompletionsHandler() key = handler.get_streaming_scan_key([self._chunk("hi"), b"data: [DONE]"]) assert key.texts == ("hi",) + + def test_released_stream_as_ended_finishes_only_the_choice_whose_tool_call_was_in_flight(self): + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + tool_call = {"index": 0, "id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}} + released = ( + self._chunk("hi", index=0), + ModelResponseStream(choices=[StreamingChoices(index=1, delta=Delta(tool_calls=[tool_call]))]), + ) + handler = OpenAIChatCompletionsHandler() + ended = handler.released_stream_as_ended(released) + assert all(a is b for a, b in zip(ended[:-1], released, strict=True)) + assert [(choice.index, choice.finish_reason) for choice in ended[-1].choices] == [(1, "tool_calls")] + ended_key = handler.get_streaming_scan_key(ended) + assert ended_key.stream_ended is True and len(ended_key.tool_calls) == 1, ended_key + + def test_released_stream_as_ended_leaves_a_stream_with_no_tool_call_in_flight_as_released(self): + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + tool_call = {"index": 0, "id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}} + text_only = (self._chunk("hi", index=0), self._chunk(" there", index=1)) + finished_tool_call = ( + ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(tool_calls=[tool_call]))]), + self._chunk(None, finish_reason="tool_calls", index=0), + ) + handler = OpenAIChatCompletionsHandler() + for released in (text_only, finished_tool_call): + ended = handler.released_stream_as_ended(released) + assert len(ended) == len(released) and all(a is b for a, b in zip(ended, released, strict=True)) diff --git a/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py b/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py index 87980a47f87..8bd242308f7 100644 --- a/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py +++ b/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py @@ -6,6 +6,7 @@ with guardrail transformations. """ import copy +import json from collections.abc import Callable from typing import Any, Final, List, Literal, Optional, Tuple from unittest.mock import AsyncMock, MagicMock, patch @@ -3654,3 +3655,85 @@ class TestOpenAIResponsesHandlerStreamingScanKey: ended_key = handler.get_streaming_scan_key([self._delta(0, "hi"), added, self._completed(3, [function_call])]) assert ended_key.tool_calls_in_flight is False assert len(ended_key.tool_calls) == 1 + + def test_released_stream_as_ended_keys_the_tool_call_the_client_already_received(self): + handler = OpenAIResponsesHandler() + added = { + "type": "response.output_item.added", + "sequence_number": 1, + "item": {"type": "function_call", "id": "fc_1", "call_id": "call_1", "name": "get_weather", "arguments": ""}, + } + arguments_delta = { + "type": "response.function_call_arguments.delta", + "sequence_number": 2, + "item_id": "fc_1", + "delta": '{"city": "Paris"', + } + ended_key = handler.get_streaming_scan_key( + handler.released_stream_as_ended([self._delta(0, "hi"), added, arguments_delta]) + ) + assert ended_key.stream_ended is True + assert ended_key.texts == ("hi",) + assert len(ended_key.tool_calls) == 1 and "Paris" in ended_key.tool_calls[0], ended_key + + def test_released_stream_as_ended_leaves_a_text_only_stream_as_released(self): + released = (self._delta(0, "hi"), self._delta(1, " there")) + ended = OpenAIResponsesHandler().released_stream_as_ended(released) + assert ended == released and all(a is b for a, b in zip(ended, released, strict=True)) + + @staticmethod + def _finished_function_call(sequence_number: int, item_id: str, city: str) -> tuple[dict[str, object], ...]: + arguments = json.dumps({"city": city}) + pending = {"type": "function_call", "id": item_id, "call_id": "call_" + item_id, "name": "get_weather"} + return ( + {"type": "response.output_item.added", "sequence_number": sequence_number, "item": {**pending, "arguments": ""}}, + { + "type": "response.function_call_arguments.delta", + "sequence_number": sequence_number + 1, + "item_id": item_id, + "delta": arguments, + }, + { + "type": "response.output_item.done", + "sequence_number": sequence_number + 2, + "item": {**pending, "arguments": arguments, "status": "completed"}, + }, + ) + + def test_released_stream_as_ended_keys_every_tool_call_finished_before_the_disconnect(self): + handler = OpenAIResponsesHandler() + released = ( + self._delta(0, "hi"), + *self._finished_function_call(1, "fc_1", "Paris"), + *self._finished_function_call(4, "fc_2", "Rome"), + ) + ended_key = handler.get_streaming_scan_key(handler.released_stream_as_ended(released)) + assert ended_key.stream_ended is True + assert ended_key.texts == ("hi",) + cities = tuple(city for fingerprint in ended_key.tool_calls for city in ("Paris", "Rome") if city in fingerprint) + assert cities == ("Paris", "Rome"), ended_key + + def test_released_stream_as_ended_keys_a_message_whose_item_already_finished(self): + handler = OpenAIResponsesHandler() + message_done = { + "type": "response.output_item.done", + "sequence_number": 1, + "item": {"type": "message", "id": "msg_1", "content": [{"type": "output_text", "text": "hi"}]}, + } + ended_key = handler.get_streaming_scan_key(handler.released_stream_as_ended((self._delta(0, "hi"), message_done))) + assert ended_key.stream_ended is True + assert ended_key.texts == ("hi",) + + @pytest.mark.asyncio + async def test_scan_of_a_stream_released_through_a_finished_tool_call_covers_its_text_too(self): + handler = OpenAIResponsesHandler() + guardrail = MockRecordingGuardrail(guardrail_name="test") + released = (self._delta(0, "hi"), *self._finished_function_call(1, "fc_1", "Paris")) + await handler.process_output_streaming_response( + responses_so_far=list(handler.released_stream_as_ended(released)), + guardrail_to_apply=guardrail, + request_data={}, + ) + assert [inputs.get("texts") for inputs in guardrail.seen_inputs] == [["hi"]], guardrail.seen_inputs + tool_calls = guardrail.seen_inputs[0].get("tool_calls") or [] + assert [call["function"]["arguments"] for call in tool_calls] == ['{"city": "Paris"}'], tool_calls diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index 45ad336368c..0cae9344160 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -5707,6 +5707,41 @@ async def test_unbuffered_end_of_stream_hook_yields_chunks_before_scan(): assert len(chunk_events) == 3 +@pytest.mark.asyncio +async def test_unbuffered_end_of_stream_hook_scans_released_chunks_when_the_client_closes_early(): + guardrail = BedrockGuardrail( + guardrail_name="bedrock-audit-mode", + guardrailIdentifier="test-id", + guardrailVersion="DRAFT", + event_hook=GuardrailEventHooks.post_call, + default_on=True, + streaming_buffer_until_moderated=False, + streaming_end_of_stream_only=True, + ) + scans = [] + + async def record_scan(*args, **kwargs): + scans.append(kwargs["source"]) + return {"action": "NONE", "assessments": [], "outputs": []} + + async def mock_stream(): + yield _chat_chunk("Hello", None) + yield _chat_chunk(" world", None) + yield _chat_chunk("", "stop") + + with patch.object(guardrail, "make_bedrock_api_request", AsyncMock(side_effect=record_scan)): + stream = guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(), + response=mock_stream(), + request_data={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]}, + ) + first = await stream.__anext__() + await stream.aclose() + + assert first.choices[0].delta.content == "Hello" + assert scans == ["OUTPUT"] + + @pytest.mark.asyncio async def test_buffered_default_hook_scans_before_any_chunk(): guardrail = BedrockGuardrail( diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index e547575ef9c..e797304a1c6 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -1,16 +1,21 @@ """Tests for unified guardrail.""" +import asyncio +import contextlib import io import logging +from collections.abc import AsyncGenerator, AsyncIterable, AsyncIterator from types import SimpleNamespace from typing import TYPE_CHECKING, Final, Literal +import anyio import pytest import litellm from litellm.caching import DualCache from litellm.integrations.custom_guardrail import ( CustomGuardrail, + ModifyResponseException, log_guardrail_information, ) from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route @@ -42,7 +47,15 @@ from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrai ) from litellm.types.guardrails import GuardrailEventHooks from litellm.types.llms.openai import ResponsesAPIResponse -from litellm.types.utils import CallTypes, Delta, GenericGuardrailAPIInputs, ModelResponseStream, StreamingChoices +from litellm.types.utils import ( + CallTypes, + Delta, + GenericGuardrailAPIInputs, + ModelResponse, + ModelResponseStream, + StandardLoggingGuardrailInformation, + StreamingChoices, +) if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -2164,6 +2177,554 @@ class _ScanCountingGuardrail(CustomGuardrail): return inputs +class _GatedScanGuardrail(_ScanCountingGuardrail): + """End-of-stream scan that holds until released, recording scans that finished.""" + + def __init__(self) -> None: + super().__init__(end_of_stream_only=True) + self.scan_started = anyio.Event() + self.scan_released = anyio.Event() + self.finished_scans = 0 + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + self.scan_started.set() + await self.scan_released.wait() + recorded = await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + self.finished_scans += 1 + return recorded + + +class _FinishReasonRecordingGuardrail(_ScanCountingGuardrail): + """Records the finish reasons of the stream handed to each response-side scan""" + + def __init__(self) -> None: + super().__init__(end_of_stream_only=True) + self.finish_reasons: tuple[str | None, ...] = () + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + rebuilt = request_data.get("response") + released = request_data.get("responses") + chunks = ( + tuple(chunk for chunk in released if isinstance(chunk, ModelResponseStream)) + if isinstance(released, list) + else () + ) + choices = [ + *(rebuilt.choices if isinstance(rebuilt, ModelResponse) else ()), + *(choice for chunk in chunks for choice in chunk.choices), + ] + self.finish_reasons = (*self.finish_reasons, *(choice.finish_reason for choice in choices)) + return await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + + +class _GatedToolCallGuardrail(_StreamingTextGuardrail): + """Tool-call inspection that holds until released""" + + def __init__(self) -> None: + super().__init__() + self.inspection_started = anyio.Event() + self.inspection_released = anyio.Event() + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + if input_type == "response" and inputs.get("tool_calls"): + self.inspection_started.set() + await self.inspection_released.wait() + return await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + + +class _MarkerBlockingScanGuardrail(_ScanCountingGuardrail): + """Scan-counting guardrail that blocks any scan whose text contains BLOCKME""" + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + recorded = await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + if any("BLOCKME" in text for text in recorded.get("texts") or []): + raise ModifyResponseException( + message="blocked", model="gpt-4", request_data=request_data, guardrail_name=self.guardrail_name + ) + return recorded + + +class _MarkerHttpErrorScanGuardrail(_ScanCountingGuardrail): + """Scan-counting guardrail that raises an HTTPException for any scan whose text contains BLOCKME""" + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + recorded = await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + if any("BLOCKME" in text for text in recorded.get("texts") or []): + raise unified_module.HTTPException(status_code=400, detail={"error": "Violated guardrail policy"}) + return recorded + + +class _MarkerBlockingStreamingTextGuardrail(_StreamingTextGuardrail): + """incremental_diff guardrail that blocks any round whose text contains BLOCKME""" + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + transformed = await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + if any("BLOCKME" in text for text in inputs.get("texts") or []): + raise ModifyResponseException( + message="blocked", model="gpt-4", request_data=request_data, guardrail_name=self.guardrail_name + ) + return transformed + + +class _DisconnectRewritingGuardrail(_ScanCountingGuardrail): + """End-of-stream guardrail that rewrites every scanned text""" + + def __init__(self) -> None: + super().__init__(end_of_stream_only=True) + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + recorded = await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + return {**recorded, "texts": ["REWRITTEN" for _ in recorded.get("texts") or []]} + + +class _RecordedScanGuardrail(CustomGuardrail): + """End-of-stream scan recorded through log_guardrail_information, returning ``reply`` or raising ``error``""" + + def __init__(self, *, reply: GenericGuardrailAPIInputs | None = None, error: Exception | None = None) -> None: + super().__init__(guardrail_name="recorded-scan") + self.streaming_end_of_stream_only = True + self.streaming_buffer_until_moderated = False + self.guardrail_config = {} + self._reply = reply + self._error = error + + def should_run_guardrail(self, data: dict[str, object], event_type: GuardrailEventHooks) -> bool: + return True + + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + if self._error is not None: + raise self._error + return inputs if self._reply is None else self._reply + + +def _recorded_guardrail_statuses(request_data: dict[str, object]) -> list[str]: + metadata = request_data["metadata"] + assert isinstance(metadata, dict), request_data + entries: list[StandardLoggingGuardrailInformation] = metadata.get("standard_logging_guardrail_information", []) + return [entry["guardrail_status"] for entry in entries] + + +class TestStreamingClientDisconnectScan: + """A client that reads streamed content and then disconnects must not skip + the end-of-stream scan of what it already received.""" + + @pytest.fixture(autouse=True) + def _use_real_mappings(self, monkeypatch: pytest.MonkeyPatch) -> None: + _patch_translation_mappings(monkeypatch, load_guardrail_translation_mappings()) + + @staticmethod + def _guarded_stream(guardrail: CustomGuardrail, upstream: AsyncIterable[object]) -> AsyncGenerator[object, None]: + return UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/chat/completions"), + response=upstream, + request_data={"guardrail_to_apply": guardrail, "model": "gpt-4", "metadata": {}}, + ) + + @pytest.mark.asyncio + async def test_closing_after_released_content_still_scans_it(self): + guardrail = _ScanCountingGuardrail(end_of_stream_only=True) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + yield _stream_chunk(" tail", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + received = await stream.__anext__() + await stream.aclose() + + assert _delta_text(received) == "synthetic secret" + assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret"]], guardrail.scans + + @pytest.mark.asyncio + async def test_closing_mid_text_stream_does_not_hand_the_scan_a_tool_calls_finish(self): + guardrail = _FinishReasonRecordingGuardrail() + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + yield _stream_chunk(" tail", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + await stream.__anext__() + await stream.aclose() + + assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret"]], guardrail.scans + assert "tool_calls" not in guardrail.finish_reasons, guardrail.finish_reasons + + @pytest.mark.asyncio + async def test_upstream_cancellation_after_released_content_still_scans_it(self): + guardrail = _ScanCountingGuardrail(end_of_stream_only=True) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + raise asyncio.CancelledError() + + stream = self._guarded_stream(guardrail, upstream()) + received = await stream.__anext__() + with pytest.raises(asyncio.CancelledError): + await stream.__anext__() + + assert _delta_text(received) == "synthetic secret" + assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret"]], guardrail.scans + + @pytest.mark.asyncio + async def test_cancellation_while_the_disconnect_scan_is_in_flight_lets_it_finish(self): + guardrail = _GatedScanGuardrail() + first_chunk_received = anyio.Event() + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + await anyio.sleep_forever() + yield _stream_chunk(" tail", finish_reason="stop") + + async def consume(scope_ready: list[anyio.CancelScope]) -> None: + with anyio.CancelScope() as scope: + scope_ready.append(scope) + async with contextlib.aclosing(self._guarded_stream(guardrail, upstream())) as stream: + async for _item in stream: + first_chunk_received.set() + + scopes = [] + with anyio.fail_after(5): + async with anyio.create_task_group() as task_group: + task_group.start_soon(consume, scopes) + await first_chunk_received.wait() + scopes[0].cancel() + await guardrail.scan_started.wait() + await anyio.sleep(0) + guardrail.scan_released.set() + + assert guardrail.finished_scans == 1 + assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret"]], guardrail.scans + + @pytest.mark.asyncio + async def test_closing_after_a_sampled_scan_covered_everything_released_does_not_scan_again(self): + guardrail = _ScanCountingGuardrail(sampling_rate=1) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic") + yield _stream_chunk(" secret") + await anyio.sleep_forever() + + stream = self._guarded_stream(guardrail, upstream()) + released = [await stream.__anext__(), await stream.__anext__()] + await stream.aclose() + + scanned_texts = [scan["texts"] for scan in guardrail.scans] + assert "".join(_delta_text(chunk) for chunk in released) == "synthetic secret" + assert scanned_texts[-1] == ["synthetic secret"], scanned_texts + assert len(scanned_texts) == len({tuple(texts) for texts in scanned_texts}), scanned_texts + + @pytest.mark.asyncio + async def test_closing_before_any_content_is_released_does_not_scan(self): + guardrail = _ScanCountingGuardrail(end_of_stream_only=True, buffer_until_moderated=True) + upstream_started = anyio.Event() + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("withheld") + upstream_started.set() + await anyio.sleep_forever() + yield _stream_chunk("never", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + async with anyio.create_task_group() as task_group: + task_group.start_soon(stream.__anext__) + await upstream_started.wait() + task_group.cancel_scope.cancel() + + assert guardrail.scans == () + + @pytest.mark.asyncio + async def test_cancellation_during_end_of_stream_scan_lets_the_scan_finish(self): + guardrail = _GatedScanGuardrail() + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + yield _stream_chunk(" tail", finish_reason="stop") + + async def consume(scope_ready: list[anyio.CancelScope]) -> None: + with anyio.CancelScope() as scope: + scope_ready.append(scope) + async for _item in self._guarded_stream(guardrail, upstream()): + pass + + scopes = [] + async with anyio.create_task_group() as task_group: + task_group.start_soon(consume, scopes) + await guardrail.scan_started.wait() + scopes[0].cancel() + await anyio.sleep(0) + guardrail.scan_released.set() + + assert guardrail.finished_scans == 1 + assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret tail"]], guardrail.scans + + @staticmethod + async def _close_after_first_chunk(guardrail: CustomGuardrail) -> dict[str, object]: + request_data: dict[str, object] = {"guardrail_to_apply": guardrail, "model": "gpt-4", "metadata": {}} + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + yield _stream_chunk(" tail", finish_reason="stop") + + stream = UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/chat/completions"), + response=upstream(), + request_data=request_data, + ) + received = await stream.__anext__() + await stream.aclose() + assert _delta_text(received) == "synthetic secret" + return request_data + + @pytest.mark.asyncio + async def test_disconnect_scan_that_fails_after_the_verdict_records_the_failure(self): + request_data = await self._close_after_first_chunk( + _RecordedScanGuardrail(reply={"texts": ["synthetic secret", "unmatched extra text"]}) + ) + + assert _recorded_guardrail_statuses(request_data) == ["success", "guardrail_failed_to_respond"] + + @pytest.mark.asyncio + async def test_disconnect_scan_whose_guardrail_raises_records_one_failure(self): + request_data = await self._close_after_first_chunk(_RecordedScanGuardrail(error=RuntimeError("provider down"))) + + assert _recorded_guardrail_statuses(request_data) == ["guardrail_failed_to_respond"] + + @pytest.mark.asyncio + async def test_disconnect_scan_that_passes_records_only_the_verdict(self): + request_data = await self._close_after_first_chunk(_RecordedScanGuardrail()) + + assert _recorded_guardrail_statuses(request_data) == ["success"] + + @staticmethod + async def _tool_call_upstream() -> AsyncIterator[ModelResponseStream]: + from litellm.types.utils import ChatCompletionDeltaToolCall, Function + + tool_call = ChatCompletionDeltaToolCall( + id="call_1", index=0, type="function", function=Function(name="get_weather", arguments='{"city": "Paris"}') + ) + yield ModelResponseStream( + choices=[StreamingChoices(index=0, delta=Delta(content=None, tool_calls=[tool_call]))] + ) + yield _stream_chunk(None, finish_reason="tool_calls") + + @pytest.mark.asyncio + async def test_cancellation_during_incremental_diff_tool_call_inspection_lets_it_finish_once(self): + guardrail = _GatedToolCallGuardrail() + + async def consume(scope_ready: list[anyio.CancelScope]) -> None: + with anyio.CancelScope() as scope: + scope_ready.append(scope) + async with contextlib.aclosing(self._guarded_stream(guardrail, self._tool_call_upstream())) as stream: + async for _item in stream: + pass + + scopes = [] + async with anyio.create_task_group() as task_group: + task_group.start_soon(consume, scopes) + await guardrail.inspection_started.wait() + scopes[0].cancel() + await anyio.sleep(0) + guardrail.inspection_released.set() + + assert [[call["function"]["name"] for call in calls] for calls in guardrail.received_tool_calls] == [ + ["get_weather"] + ], guardrail.received_tool_calls + + @pytest.mark.asyncio + async def test_closing_after_the_incremental_diff_tool_call_inspection_does_not_inspect_again(self): + guardrail = _StreamingTextGuardrail() + + stream = self._guarded_stream(guardrail, self._tool_call_upstream()) + released = [await stream.__anext__(), await stream.__anext__()] + await stream.aclose() + + assert released[-1].choices[0].finish_reason == "tool_calls" + assert [[call["function"]["name"] for call in calls] for calls in guardrail.received_tool_calls] == [ + ["get_weather"] + ], guardrail.received_tool_calls + + @pytest.mark.asyncio + async def test_closing_after_a_released_tool_call_under_incremental_diff_still_inspects_it(self): + guardrail = _StreamingTextGuardrail() + + stream = self._guarded_stream(guardrail, self._tool_call_upstream()) + received = await stream.__anext__() + await stream.aclose() + + assert [call.function.name for call in received.choices[0].delta.tool_calls] == ["get_weather"] + assert [[call["function"]["name"] for call in calls] for calls in guardrail.received_tool_calls] == [ + ["get_weather"] + ], guardrail.received_tool_calls + + @pytest.mark.asyncio + async def test_closing_while_a_mid_stream_block_is_delivered_does_not_scan_the_blocked_content_again(self): + guardrail = _MarkerBlockingScanGuardrail(sampling_rate=2) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + for text in ("a", "b", "c", "BLOCKME"): + yield _stream_chunk(text) + yield _stream_chunk(" tail", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + received = [_delta_text(await stream.__anext__()) for _ in range(3)] + await stream.__anext__() + await stream.aclose() + + assert received == ["a", "b", "c"] + assert [scan["texts"] for scan in guardrail.scans] == [["ab"], ["abcBLOCKME"]], guardrail.scans + + @pytest.mark.asyncio + async def test_closing_while_a_mid_stream_guardrail_error_is_delivered_does_not_scan_the_content_again(self): + guardrail = _MarkerHttpErrorScanGuardrail(sampling_rate=2) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + for text in ("a", "b", "c", "BLOCKME"): + yield _stream_chunk(text) + yield _stream_chunk(" tail", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + received = [_delta_text(await stream.__anext__()) for _ in range(3)] + await stream.__anext__() + await stream.aclose() + + assert received == ["a", "b", "c"] + assert [scan["texts"] for scan in guardrail.scans] == [["ab"], ["abcBLOCKME"]], guardrail.scans + + @pytest.mark.asyncio + async def test_closing_while_the_scanned_final_chunk_is_delivered_does_not_scan_the_stream_again(self): + guardrail = _ScanCountingGuardrail(end_of_stream_only=True) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("a") + yield _stream_chunk("b") + yield _stream_chunk(" tail", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + received = [_delta_text(await stream.__anext__()) for _ in range(3)] + await stream.aclose() + + assert received == ["a", "b", " tail"] + assert [scan["texts"] for scan in guardrail.scans] == [["ab tail"]], guardrail.scans + + @pytest.mark.asyncio + async def test_cancellation_with_a_withheld_window_scans_only_the_released_chunks(self): + guardrail = _ScanCountingGuardrail(sampling_rate=2, buffer_until_moderated=True) + guardrail.streaming_buffer_release_on_scan = True + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("a") + yield _stream_chunk("b") + yield _stream_chunk("WITHHELD") + raise asyncio.CancelledError() + + stream = self._guarded_stream(guardrail, upstream()) + received = [_delta_text(await stream.__anext__()), _delta_text(await stream.__anext__())] + with pytest.raises(asyncio.CancelledError): + await stream.__anext__() + + assert received == ["a", "b"] + assert [scan["texts"] for scan in guardrail.scans] == [["ab"]], guardrail.scans + + @pytest.mark.asyncio + async def test_disconnect_scan_that_rewrites_text_leaves_the_released_chunks_untouched(self): + stream = self._guarded_stream(_DisconnectRewritingGuardrail(), self._text_then_tail_upstream()) + received = await stream.__anext__() + await stream.aclose() + + assert _delta_text(received) == "synthetic secret" + + @staticmethod + async def _text_then_tail_upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + yield _stream_chunk(" tail", finish_reason="stop") + + @pytest.mark.asyncio + async def test_closing_while_an_incremental_diff_block_is_delivered_does_not_inspect_the_tool_call_again(self): + guardrail = _MarkerBlockingStreamingTextGuardrail() + + async def upstream() -> AsyncIterator[ModelResponseStream]: + async for tool_chunk in self._tool_call_upstream(): + if tool_chunk.choices[0].finish_reason is None: + yield tool_chunk + yield _stream_chunk("BLOCKME") + yield _stream_chunk(" tail", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + released = [await stream.__anext__(), await stream.__anext__()] + await stream.aclose() + + assert [call.function.name for call in released[0].choices[0].delta.tool_calls] == ["get_weather"] + assert guardrail.received_texts == [["BLOCKME"]], guardrail.received_texts + assert guardrail.received_tool_calls == [], guardrail.received_tool_calls + + @pytest.mark.asyncio + async def test_closing_after_a_released_tool_call_under_incremental_diff_scans_no_held_back_text(self): + guardrail = _StreamingTextGuardrail(holdback_schedule=[len("held secret")] * 2) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("held secret") + async for tool_chunk in self._tool_call_upstream(): + yield tool_chunk + + stream = self._guarded_stream(guardrail, upstream()) + received = await stream.__anext__() + await stream.aclose() + + assert [call.function.name for call in received.choices[0].delta.tool_calls] == ["get_weather"] + assert guardrail.received_tool_calls, guardrail.received_texts + assert all("held secret" not in text for text in guardrail.received_texts[-1]), guardrail.received_texts + + def _responses_delta(sequence_number, text): return { "type": "response.output_text.delta", diff --git a/tests/unit/proxy/hooks/test_async_post_call_streaming_iterator_hook.py b/tests/unit/proxy/hooks/test_async_post_call_streaming_iterator_hook.py index 9c785d59830..31f1c054c59 100644 --- a/tests/unit/proxy/hooks/test_async_post_call_streaming_iterator_hook.py +++ b/tests/unit/proxy/hooks/test_async_post_call_streaming_iterator_hook.py @@ -7,6 +7,7 @@ Verifies that the hook: 3. Actually yields chunks from async generators """ +import logging from typing import AsyncGenerator, Any from unittest.mock import MagicMock, patch @@ -31,7 +32,7 @@ class MockStreamingCallback(CustomLogger): self, user_api_key_dict: UserAPIKeyAuth, response: AsyncGenerator[Any, None], - request_data: dict, + request_data: dict[str, object], ) -> AsyncGenerator[Any, None]: """Transform chunks by tracking and optionally prefixing.""" async for chunk in response: @@ -185,3 +186,340 @@ async def test_streaming_hook_propagates_callback_errors(): with pytest.raises(RuntimeError, match="Callback failed!"): async for _ in result: pass + + +class CleanupRecordingCallback(CustomLogger): + """Iterator hook whose cleanup marks when it ran.""" + + def __init__(self): + super().__init__() + self.cleaned_up = False + + async def async_post_call_streaming_iterator_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncGenerator[Any, None], + request_data: dict[str, object], + ) -> AsyncGenerator[Any, None]: + try: + async for chunk in response: + yield chunk + finally: + self.cleaned_up = True + + +@pytest.mark.asyncio +async def test_closing_the_stream_runs_every_callback_cleanup_before_returning(): + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + callbacks = [CleanupRecordingCallback(), CleanupRecordingCallback()] + + with patch.object(litellm, "callbacks", callbacks): + ProxyLogging._callback_capabilities_cache.clear() + stream = proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + first = await stream.__anext__() + await stream.aclose() + ProxyLogging._callback_capabilities_cache.clear() + + assert first == {"choices": [{"delta": {"content": "Hello"}}]} + assert [callback.cleaned_up for callback in callbacks] == [True, True] + + +class RaisingCleanupCallback(CustomLogger): + """Iterator hook whose cleanup raises.""" + + async def async_post_call_streaming_iterator_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncGenerator[Any, None], + request_data: dict[str, object], + ) -> AsyncGenerator[Any, None]: + try: + async for chunk in response: + yield chunk + finally: + raise RuntimeError("cleanup failed") + + +@pytest.mark.asyncio +async def test_closing_the_stream_still_cleans_up_inner_callbacks_when_an_outer_cleanup_raises(): + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + inner = CleanupRecordingCallback() + + with patch.object(litellm, "callbacks", [inner, RaisingCleanupCallback()]): + ProxyLogging._callback_capabilities_cache.clear() + stream = proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + await stream.__anext__() + await stream.aclose() + ProxyLogging._callback_capabilities_cache.clear() + + assert inner.cleaned_up is True + + +class _PlainAsyncIterator: + def __init__(self, response: AsyncGenerator[Any, None]) -> None: + self._response = response + + def __aiter__(self) -> "_PlainAsyncIterator": + return self + + async def __anext__(self) -> Any: + return await self._response.__anext__() + + +class PlainIteratorCallback(CustomLogger): + """Iterator hook that returns an async iterator with no aclose.""" + + def async_post_call_streaming_iterator_hook( # pyright: ignore[reportIncompatibleMethodOverride] # a plain async iterator worked before aclose handling + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncGenerator[Any, None], + request_data: dict[str, object], + ) -> _PlainAsyncIterator: + return _PlainAsyncIterator(response) + + +@pytest.mark.asyncio +async def test_a_hook_returning_a_plain_async_iterator_streams_every_chunk(): + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + + with patch.object(litellm, "callbacks", [PlainIteratorCallback()]): + ProxyLogging._callback_capabilities_cache.clear() + received = [ + chunk + async for chunk in proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + ] + ProxyLogging._callback_capabilities_cache.clear() + + assert [chunk async for chunk in mock_streaming_response()] == received + + +class _ClosableAsyncIterator(_PlainAsyncIterator): + def __init__(self, response: AsyncGenerator[Any, None]) -> None: + super().__init__(response) + self.closed = False + + async def aclose(self) -> None: + self.closed = True + + +class _SyncClosableAsyncIterator(_PlainAsyncIterator): + def __init__(self, response: AsyncGenerator[Any, None]) -> None: + super().__init__(response) + self.closed = False + + def aclose(self) -> None: + self.closed = True + + +class ClosableIteratorCallback(CustomLogger): + """Iterator hook that returns a non-generator async iterator with its own aclose.""" + + def __init__( + self, + iterator_type: type[_ClosableAsyncIterator] | type[_SyncClosableAsyncIterator] = _ClosableAsyncIterator, + ) -> None: + super().__init__() + self.iterator_type = iterator_type + self.returned: tuple[_ClosableAsyncIterator | _SyncClosableAsyncIterator, ...] = () + + def async_post_call_streaming_iterator_hook( # pyright: ignore[reportIncompatibleMethodOverride] # a custom async iterator worked before aclose handling + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncGenerator[Any, None], + request_data: dict[str, object], + ) -> _ClosableAsyncIterator | _SyncClosableAsyncIterator: + iterator = self.iterator_type(response) + self.returned = (*self.returned, iterator) + return iterator + + +@pytest.mark.asyncio +async def test_closing_the_stream_closes_a_hook_iterator_that_is_not_a_generator(): + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + callback = ClosableIteratorCallback() + + with patch.object(litellm, "callbacks", [callback]): + ProxyLogging._callback_capabilities_cache.clear() + stream = proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + await stream.__anext__() + await stream.aclose() + ProxyLogging._callback_capabilities_cache.clear() + + assert [iterator.closed for iterator in callback.returned] == [True] + + +@pytest.mark.asyncio +async def test_a_hook_iterator_with_a_synchronous_aclose_streams_everything_and_is_closed(): + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + callback = ClosableIteratorCallback(iterator_type=_SyncClosableAsyncIterator) + + with patch.object(litellm, "callbacks", [callback]): + ProxyLogging._callback_capabilities_cache.clear() + received = [ + chunk + async for chunk in proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + ] + ProxyLogging._callback_capabilities_cache.clear() + + assert [chunk async for chunk in mock_streaming_response()] == received + assert [iterator.closed for iterator in callback.returned] == [True] + + +class _RaisingAcloseIterator(_ClosableAsyncIterator): + """Non-generator async iterator whose asynchronous aclose raises.""" + + def __init__(self, response: AsyncGenerator[Any, None], error: Exception) -> None: + super().__init__(response) + self.error = error + + async def aclose(self) -> None: + self.closed = True + raise self.error + + +class _SyncRaisingAcloseIterator(_SyncClosableAsyncIterator): + """Non-generator async iterator whose synchronous aclose raises.""" + + def __init__(self, response: AsyncGenerator[Any, None], error: Exception) -> None: + super().__init__(response) + self.error = error + + def aclose(self) -> None: + self.closed = True + raise self.error + + +class RaisingAcloseIteratorCallback(CustomLogger): + """Iterator hook that returns a non-generator async iterator whose aclose raises.""" + + def __init__( + self, + iterator_type: type[_RaisingAcloseIterator] | type[_SyncRaisingAcloseIterator] = _RaisingAcloseIterator, + error: Exception | None = None, + ) -> None: + super().__init__() + self.iterator_type = iterator_type + self.error = error if error is not None else RuntimeError("cleanup failed") + self.returned: tuple[_RaisingAcloseIterator | _SyncRaisingAcloseIterator, ...] = () + + def async_post_call_streaming_iterator_hook( # pyright: ignore[reportIncompatibleMethodOverride] # a custom async iterator worked before aclose handling + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncGenerator[Any, None], + request_data: dict[str, object], + ) -> _RaisingAcloseIterator | _SyncRaisingAcloseIterator: + iterator = self.iterator_type(response, self.error) + self.returned = (*self.returned, iterator) + return iterator + + +@pytest.mark.asyncio +async def test_a_hook_iterator_whose_aclose_raises_still_finishes_the_stream(caplog: pytest.LogCaptureFixture) -> None: + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + callback = RaisingAcloseIteratorCallback() + + with patch.object(litellm, "callbacks", [callback]): + ProxyLogging._callback_capabilities_cache.clear() + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + received = [ + chunk + async for chunk in proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + ] + ProxyLogging._callback_capabilities_cache.clear() + + assert [chunk async for chunk in mock_streaming_response()] == received + assert [iterator.closed for iterator in callback.returned] == [True] + warnings_emitted = [ + record.getMessage() + for record in caplog.records + if record.levelname == "WARNING" and "RaisingAcloseIteratorCallback" in record.getMessage() + ] + assert len(warnings_emitted) == 1 + assert "RuntimeError" in warnings_emitted[0] + assert "cleanup failed" not in warnings_emitted[0] + + +@pytest.mark.asyncio +async def test_a_hook_iterator_whose_synchronous_aclose_raises_still_finishes_the_stream( + caplog: pytest.LogCaptureFixture, +) -> None: + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + callback = RaisingAcloseIteratorCallback( + iterator_type=_SyncRaisingAcloseIterator, error=ValueError("sync cleanup failed") + ) + + with patch.object(litellm, "callbacks", [callback]): + ProxyLogging._callback_capabilities_cache.clear() + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + received = [ + chunk + async for chunk in proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + ] + ProxyLogging._callback_capabilities_cache.clear() + + assert [chunk async for chunk in mock_streaming_response()] == received + assert [iterator.closed for iterator in callback.returned] == [True] + warnings_emitted = [ + record.getMessage() + for record in caplog.records + if record.levelname == "WARNING" and "RaisingAcloseIteratorCallback" in record.getMessage() + ] + assert len(warnings_emitted) == 1 + assert "ValueError" in warnings_emitted[0] + + +@pytest.mark.asyncio +async def test_a_hook_iterator_with_a_clean_aclose_streams_everything_without_warning( + caplog: pytest.LogCaptureFixture, +) -> None: + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + callback = ClosableIteratorCallback() + + with patch.object(litellm, "callbacks", [callback]): + ProxyLogging._callback_capabilities_cache.clear() + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + received = [ + chunk + async for chunk in proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + ] + ProxyLogging._callback_capabilities_cache.clear() + + assert [chunk async for chunk in mock_streaming_response()] == received + assert [iterator.closed for iterator in callback.returned] == [True] + assert not [ + record.getMessage() + for record in caplog.records + if record.levelname == "WARNING" and "ClosableIteratorCallback" in record.getMessage() + ] diff --git a/tests/unit/proxy/test_common_request_processing.py b/tests/unit/proxy/test_common_request_processing.py index 8485c286a30..8dc6c0477b0 100644 --- a/tests/unit/proxy/test_common_request_processing.py +++ b/tests/unit/proxy/test_common_request_processing.py @@ -26,6 +26,7 @@ from litellm.constants import ( CLIENT_REQUESTED_MODEL_SCOPE_KEY, MAX_LITELLM_CALL_ID_LENGTH, RETURN_RAW_MODEL_NAME_METADATA_KEY, + STREAM_SSE_KEEPALIVE_PING_BYTES, ) from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.opentelemetry import UserAPIKeyAuth @@ -4256,6 +4257,72 @@ class TestStreamCloseOnDisconnect: assert upstream.aclosed + async def test_async_streaming_data_generator_closes_the_guardrail_chain_on_client_disconnect( + self, + ): + cleanup_ran = [] + + async def guarded_chain(**_kwargs): + try: + yield {"type": "chunk"} + yield {"type": "chunk"} + finally: + cleanup_ran.append(True) + + proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock()) + proxy_logging_obj.async_post_call_streaming_iterator_hook = guarded_chain + gen = ProxyBaseLLMRequestProcessing.async_streaming_data_generator( + response=MagicMock(), + user_api_key_dict=MagicMock(spec=UserAPIKeyAuth), + request_data={"model": "mock-model"}, + proxy_logging_obj=proxy_logging_obj, + serialize_chunk=lambda c: "data: x\n\n", + serialize_error=lambda e: "data: error\n\n", + ) + + await gen.__anext__() + await gen.aclose() + + assert cleanup_ran == [True] + + async def test_async_streaming_data_generator_refunds_the_budget_when_closing_the_guardrail_chain_raises( + self, + ): + async def guarded_chain(**_kwargs): + try: + yield STREAM_SSE_KEEPALIVE_PING_BYTES + yield {"type": "chunk"} + finally: + raise RuntimeError("cleanup failed") + + proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock()) + proxy_logging_obj.async_post_call_streaming_iterator_hook = guarded_chain + reservation = object() + user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + user_api_key_dict.budget_reservation = reservation + gen = ProxyBaseLLMRequestProcessing.async_streaming_data_generator( + response=MagicMock(), + user_api_key_dict=user_api_key_dict, + request_data={"model": "mock-model"}, + proxy_logging_obj=proxy_logging_obj, + serialize_chunk=lambda c: "data: x\n\n", + serialize_error=lambda e: "data: error\n\n", + ) + + released: list[object] = [] + + async def record_release(budget_reservation: object) -> None: + released.append(budget_reservation) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation_on_cancel", + new=record_release, + ): + await gen.__anext__() + await gen.aclose() + + assert released == [reservation] + async def test_async_streaming_data_generator_redacts_internal_details_on_error( self, ): diff --git a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py index 2397c6480d5..e2fbd5dcd98 100644 --- a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py +++ b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py @@ -7425,6 +7425,69 @@ async def test_async_data_generator_cleanup_on_early_exit(): mock_response.aclose.assert_awaited_once() +def _guarded_chain_logging(chain): + from litellm.proxy.utils import ProxyLogging + + proxy_logging = MagicMock(spec=ProxyLogging) + proxy_logging.async_post_call_streaming_iterator_hook = chain + proxy_logging.async_post_call_streaming_hook = AsyncMock(side_effect=lambda **kwargs: kwargs.get("response")) + proxy_logging.post_call_failure_hook = AsyncMock() + return proxy_logging + + +@pytest.mark.asyncio +async def test_async_data_generator_closes_the_guardrail_chain_before_returning_on_client_disconnect(): + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.proxy_server import async_data_generator + + cleanup_ran = [] + + async def guarded_chain(**_kwargs): + try: + yield {"choices": [{"delta": {"content": "Hello"}}]} + yield {"choices": [{"delta": {"content": " world"}}]} + finally: + cleanup_ran.append(True) + + with patch("litellm.proxy.proxy_server.proxy_logging_obj", _guarded_chain_logging(guarded_chain)): + gen = async_data_generator(MagicMock(), MagicMock(spec=UserAPIKeyAuth), {"model": "gpt-4o-mini"}) + first_chunk = await gen.__anext__() + await gen.aclose() + + assert first_chunk.startswith("data: ") + assert cleanup_ran == [True] + + +@pytest.mark.asyncio +async def test_async_data_generator_closes_the_guardrail_chain_while_a_keepalive_read_is_pending(): + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.proxy_server import async_data_generator + + cleanup_ran = [] + never_arrives = asyncio.Event() + + async def guarded_chain(**_kwargs): + try: + yield {"choices": [{"delta": {"content": "Hello"}}]} + await never_arrives.wait() + yield {"choices": [{"delta": {"content": " world"}}]} + finally: + cleanup_ran.append(True) + + with ( + patch.object(litellm, "sse_keepalive_ping_interval_seconds", 1.0), + patch("litellm.proxy.proxy_server.proxy_logging_obj", _guarded_chain_logging(guarded_chain)), + ): + gen = async_data_generator(MagicMock(), MagicMock(spec=UserAPIKeyAuth), {"model": "gpt-4o-mini"}) + first_chunk = await gen.__anext__() + heartbeat = await gen.__anext__() + await gen.aclose() + + assert first_chunk.startswith("data: ") + assert heartbeat == ": ping\n\n" + assert cleanup_ran == [True] + + @pytest.mark.asyncio async def test_async_data_generator_uses_direct_stream_fast_path_without_callbacks(): """ diff --git a/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py b/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py index 8bd9dc0df8a..dbb354526d8 100644 --- a/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py +++ b/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py @@ -11,9 +11,9 @@ from __future__ import annotations import asyncio import json +from collections.abc import AsyncGenerator, Iterator from copy import deepcopy import logging -from collections.abc import Iterator from typing import Any, Callable, Dict, List from unittest.mock import AsyncMock, MagicMock, patch @@ -26,6 +26,7 @@ from litellm.integrations.custom_guardrail import ( CustomGuardrail, ModifyResponseException, ) +from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.prometheus import PrometheusLogger from litellm.llms.base_llm.guardrail_translation.utils import stream_item_field from litellm.proxy._types import UserAPIKeyAuth @@ -2933,3 +2934,62 @@ async def test_streaming_iterator_hook_pipeline_discards_dropped_tool_call_on_re assert delivered == _responses_function_call_events() assert any("'gr-post'" in message and "discarded" in message for message in _warnings(caplog)) + + +class _RaisingAcloseIterator: + """Non-generator async iterator whose aclose raises after the stream ended.""" + + def __init__(self, response: AsyncGenerator[Any, None]) -> None: + self._response = response + + def __aiter__(self) -> "_RaisingAcloseIterator": + return self + + async def __anext__(self) -> Any: + return await self._response.__anext__() + + async def aclose(self) -> None: + raise RuntimeError("cleanup failed") + + +class RaisingAcloseCallback(CustomLogger): + """Iterator hook returning a non-generator async iterator whose aclose raises.""" + + def async_post_call_streaming_iterator_hook( # pyright: ignore[reportIncompatibleMethodOverride] # a custom async iterator worked before aclose handling + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncGenerator[Any, None], + request_data: dict[str, object], + ) -> _RaisingAcloseIterator: + return _RaisingAcloseIterator(response) + + +@pytest.mark.asyncio +async def test_streaming_iterator_hook_pipeline_releases_buffered_content_when_a_callback_aclose_raises( + proxy_logging: ProxyLogging, + make_user_api_key_auth: Callable[..., UserAPIKeyAuth], + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +) -> None: + monkeypatch.setattr( + litellm, "callbacks", [_rewriting_stream_guardrail(lambda inputs: {}), RaisingAcloseCallback()] + ) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False) + data = _post_call_pipeline_data(stream=True) + chunks = _tool_call_stream_chunks() + + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + delivered = [ + item + async for item in proxy_logging.async_post_call_streaming_iterator_hook( + user_api_key_dict=make_user_api_key_auth(request_route="/v1/chat/completions"), + response=_async_chunk_iter(chunks), + request_data=data, + ) + ] + + assert [chunk.model_dump() for chunk in delivered] == [chunk.model_dump() for chunk in chunks] + assert any( + "RaisingAcloseCallback" in message and "RuntimeError" in message and "cleanup failed" not in message + for message in _warnings(caplog) + ) From 254c2f6b3b1e75c4074e33ee416480e008632daf Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 5 Oct 2026 09:07:23 +0000 Subject: [PATCH 03/19] refactor(mcp): drop Sequence/list return-type mismatch and collapse record_listed_tools wrapper (#44556) Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 34 +++++-------- .../_experimental/mcp_server/operations.py | 26 +++++----- .../mcp_server/rest_endpoints.py | 22 +++++---- .../mcp_server/test_mcp_server_manager.py | 48 +++++++++---------- .../test_mcp_server_tool_calls_and_headers.py | 10 ++-- 5 files changed, 67 insertions(+), 73 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 49cd8cf6fb9..b8099c97861 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -3801,7 +3801,7 @@ class MCPServerManager: if server is None: verbose_logger.warning("MCP Server %s not found", server_id) return [] - return await self._get_tools_from_server(server) + return list(await self._get_tools_from_server(server)) except Exception as e: verbose_logger.warning("Failed to get tools from server %s: %s", server_id, e) return [] @@ -3838,11 +3838,13 @@ class MCPServerManager: server_auth_header: Final = _server_auth_header_for(server, mcp_server_auth_headers, mcp_auth_header) try: - tools: Final = await self._get_tools_from_server( - server=server, - mcp_auth_header=server_auth_header, - user_api_key_auth=user_api_key_auth, - record_listing=True, + tools: Final = list( + await self._get_tools_from_server( + server=server, + mcp_auth_header=server_auth_header, + user_api_key_auth=user_api_key_auth, + record_listing=True, + ) ) return tools except Exception as e: @@ -4350,12 +4352,11 @@ class MCPServerManager: # applied (e.g. "test_petstore-getinventory"). Do NOT pass them # through _create_prefixed_tools — that would add the prefix a second # time producing "test_petstore-test_petstore-getinventory". - unprefixed_tools: Final = guarded_openapi - self._record_listed_tools( - server, unprefixed_tools, listed_caller, listed_generation, record_listing=record_listing + self.record_listed_tools( + server, guarded_openapi, listed_caller, listed_generation, record_listing=record_listing ) if not add_prefix: - return unprefixed_tools + return guarded_openapi return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi] else: tools = await self._fetch_tools_with_timeout(client, server.name) @@ -4371,7 +4372,7 @@ class MCPServerManager: prefixed_or_original_tools: Final = self._create_prefixed_tools( guarded_tools, server, add_prefix=add_prefix ) - self._record_listed_tools( + self.record_listed_tools( server, guarded_tools, listed_caller, listed_generation, record_listing=record_listing ) @@ -4478,17 +4479,6 @@ class MCPServerManager: return self._listed_tools_generations.get(server_id, 0) def record_listed_tools( - self, - server: MCPServer, - tools: Sequence[MCPTool], - caller: ListedToolsCaller | None, - generation: int, - *, - record_listing: bool = True, - ) -> None: - self._record_listed_tools(server, tools, caller, generation, record_listing=record_listing) - - def _record_listed_tools( self, server: MCPServer, tools: Sequence[MCPTool], diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 1dea44a84f0..42cca179cd8 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -1125,18 +1125,20 @@ async def _get_tools_from_mcp_servers( from litellm.proxy.proxy_server import proxy_logging_obj listed_generation: Final = global_mcp_server_manager.listed_tools_generation(server.server_id) - tools: Final = await global_mcp_server_manager._get_tools_from_server( - server=server, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - add_prefix=True, # Always add server prefix - raw_headers=raw_headers, - client_ip=client_ip, - user_api_key_auth=user_api_key_auth, - oauth2_headers=oauth2_headers, - proxy_logging_obj=proxy_logging_obj, - catalog_auth_header=catalog_auth_header, - record_listing=False, + tools: Final = list( + await global_mcp_server_manager._get_tools_from_server( + server=server, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=True, # Always add server prefix + raw_headers=raw_headers, + client_ip=client_ip, + user_api_key_auth=user_api_key_auth, + oauth2_headers=oauth2_headers, + proxy_logging_obj=proxy_logging_obj, + catalog_auth_header=catalog_auth_header, + record_listing=False, + ) ) filtered_tools = filter_tools_by_allowed_tools(tools, server) diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 845dcaa1f16..534eda07292 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -705,16 +705,18 @@ if MCP_AVAILABLE: *, record_listing: bool, ) -> list[MCPTool]: - return await global_mcp_server_manager._get_tools_from_server( - server=server, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - add_prefix=False, - raw_headers=raw_headers, - client_ip=client_ip, - user_api_key_auth=user_api_key_auth, - proxy_logging_obj=proxy_logging_obj, - record_listing=record_listing, + return list( + await global_mcp_server_manager._get_tools_from_server( + server=server, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=False, + raw_headers=raw_headers, + client_ip=client_ip, + user_api_key_auth=user_api_key_auth, + proxy_logging_obj=proxy_logging_obj, + record_listing=record_listing, + ) ) async def _get_tools_for_single_server( diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index cb017afbea5..e7602d7f825 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -7197,7 +7197,7 @@ class TestMCPServerManager: manager.tool_name_to_mcp_server_name_mapping["test_tool"] = "test-server" manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = "test-server" manager._create_prefixed_tools(listed_tools, server) - manager._record_listed_tools(server, listed_tools, caller) + manager.record_listed_tools(server, listed_tools, caller) mock_client = AsyncMock() mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) @@ -7281,8 +7281,8 @@ class TestMCPServerManager: def test_get_listed_tool_resolves_the_bare_name_from_the_latest_listing(self): manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") - manager._record_listed_tools(server, [MCPTool(name="echo", description="v1", inputSchema={})], None) - manager._record_listed_tools(server, [MCPTool(name="echo", description="v2", inputSchema={})], None) + manager.record_listed_tools(server, [MCPTool(name="echo", description="v1", inputSchema={})], None) + manager.record_listed_tools(server, [MCPTool(name="echo", description="v2", inputSchema={})], None) latest = manager.get_listed_tool(server, "echo") assert latest is not None and latest.description == "v2" @@ -7294,7 +7294,7 @@ class TestMCPServerManager: manager = MCPServerManager() server = MCPServer(server_id="srv-id", name="srv", alias="srv", transport=MCPTransport.http, url="http://srv") caller = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-user", user_id="alice")) - manager._record_listed_tools( + manager.record_listed_tools( server, [ MCPTool(name="foo", description="Fetches foo records", inputSchema={"type": "object"}), @@ -7355,8 +7355,8 @@ class TestMCPServerManager: manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") other = MCPServer(server_id="other", name="other", transport=MCPTransport.http, url="http://other") - manager._record_listed_tools(server, [MCPTool(name="echo", description="old", inputSchema={})], None) - manager._record_listed_tools(other, [MCPTool(name="ping", description="kept", inputSchema={})], None) + manager.record_listed_tools(server, [MCPTool(name="echo", description="old", inputSchema={})], None) + manager.record_listed_tools(other, [MCPTool(name="ping", description="kept", inputSchema={})], None) manager._invalidate_server_definition_caches(server.server_id) @@ -7418,7 +7418,7 @@ class TestMCPServerManager: caller = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="lister")) async def register_while_a_listing_records(server: MCPServer, *, initialize_mapping: bool = True) -> None: - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="search", description="pre-save", inputSchema={})], caller, @@ -7463,7 +7463,7 @@ class TestMCPServerManager: return [Prompt(name="greet")] async def register_while_discovery_fills(server: MCPServer, *, initialize_mapping: bool = True) -> None: - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="search", description="pre-save", inputSchema={})], caller, @@ -7497,7 +7497,7 @@ class TestMCPServerManager: async def test_user_oauth_refresh_keeps_listed_tools(self): manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") - manager._record_listed_tools(server, [MCPTool(name="echo", description="shared", inputSchema={})], None) + manager.record_listed_tools(server, [MCPTool(name="echo", description="shared", inputSchema={})], None) await manager.invalidate_user_oauth_token_cache("alice", server.server_id) @@ -7517,12 +7517,12 @@ class TestMCPServerManager: bob = UserAPIKeyAuth(user_id="bob", token="hashed-bob") alice_schema = {"type": "object", "properties": {"path": {"type": "string"}}} bob_schema = {"type": "object", "properties": {"path": {"type": "string"}, "site": {"type": "string"}}} - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="read", description="alice view", inputSchema=alice_schema)], ListedToolsCaller(user_api_key_auth=alice), ) - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="read", description="bob view", inputSchema=bob_schema)], ListedToolsCaller(user_api_key_auth=bob), @@ -7539,7 +7539,7 @@ class TestMCPServerManager: assert manager.get_listed_tool(server, "read", carol) is None shared = MCPServer(server_id="shared", name="shared", transport=MCPTransport.http, url="http://shared") - manager._record_listed_tools( + manager.record_listed_tools( shared, [MCPTool(name="echo", description="everyone", inputSchema={})], ListedToolsCaller(user_api_key_auth=alice), @@ -7595,8 +7595,8 @@ class TestMCPServerManager: server = MCPServer( **{"server_id": "srv", "name": "srv", "transport": MCPTransport.http, "url": "http://srv", **server_kwargs} ) - manager._record_listed_tools(server, [MCPTool(name="turn", description="Catalog A", inputSchema={})], caller_a) - manager._record_listed_tools(server, [MCPTool(name="turn", description="Catalog B", inputSchema={})], caller_b) + manager.record_listed_tools(server, [MCPTool(name="turn", description="Catalog A", inputSchema={})], caller_a) + manager.record_listed_tools(server, [MCPTool(name="turn", description="Catalog B", inputSchema={})], caller_b) for_a = manager.get_listed_tool(server, "turn", caller_a) for_b = manager.get_listed_tool(server, "turn", caller_b) @@ -7607,7 +7607,7 @@ class TestMCPServerManager: def test_shared_server_ignores_headers_it_never_forwards(self): manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="turn", description="everyone", inputSchema={})], ListedToolsCaller(raw_headers={"authorization": "Bearer sk-litellm", "x-workspace": "A"}), @@ -7840,7 +7840,7 @@ class TestMCPServerManager: "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", return_value=signer, ): - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="turn", description="alice view", inputSchema={})], alice ) assert manager.get_listed_tool(server, "turn", bob) is None @@ -7860,7 +7860,7 @@ class TestMCPServerManager: "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", return_value=MagicMock(), ): - manager._record_listed_tools(server, [MCPTool(name="turn", description="slot a", inputSchema={})], alice) + manager.record_listed_tools(server, [MCPTool(name="turn", description="slot a", inputSchema={})], alice) assert manager.get_listed_tool(server, "turn", bob) is None same_key = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-alpha")) @@ -7878,7 +7878,7 @@ class TestMCPServerManager: team_two: Final = ListedToolsCaller( user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-two") ) - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="foo", description="Fetch rows FLAGWORD", inputSchema={})], team_one ) @@ -7896,7 +7896,7 @@ class TestMCPServerManager: alice_in_two: Final = ListedToolsCaller( user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id="alice", team_id="team-two") ) - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="foo", description="Fetch rows FLAGWORD", inputSchema={})], alice_in_one ) @@ -7917,7 +7917,7 @@ class TestMCPServerManager: user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-one"), raw_headers={"authorization": "Bearer jwt-bob"}, ) - manager._record_listed_tools(server, [MCPTool(name="foo", description="alice view", inputSchema={})], alice) + manager.record_listed_tools(server, [MCPTool(name="foo", description="alice view", inputSchema={})], alice) assert manager.get_listed_tool(server, "foo", bob) is None listed: Final = manager.get_listed_tool(server, "foo", alice) @@ -7976,7 +7976,7 @@ class TestMCPServerManager: user_api_key_auth=UserAPIKeyAuth(api_key="sk-master"), raw_headers={"x-litellm-api-key": "Bearer sk-master", "authorization": "Bearer UP-B"}, ) - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="lookup", description="Workspace A lookup FLAGWORD", inputSchema={})], caller_a ) @@ -8053,18 +8053,18 @@ class TestMCPServerManager: url="http://srv", auth_type=MCPAuth.oauth2_token_exchange, ) - manager._record_listed_tools(server, [MCPTool(name="read", description="shared", inputSchema={})], None) + manager.record_listed_tools(server, [MCPTool(name="read", description="shared", inputSchema={})], None) callers = [ ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id=f"u{i}", api_key=f"k{i}")) for i in range(_LISTED_TOOLS_CALLERS_PER_SERVER + 1) ] for caller in callers: - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="read", description=caller.user_api_key_auth.user_id, inputSchema={})], caller, ) - manager._record_listed_tools(server, [MCPTool(name="read", description="u1 again", inputSchema={})], callers[1]) + manager.record_listed_tools(server, [MCPTool(name="read", description="u1 again", inputSchema={})], callers[1]) assert manager.get_listed_tool(server, "read", callers[0]) is None second = manager.get_listed_tool(server, "read", callers[1]) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index 556fe0b2536..3b53d023af3 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -8115,7 +8115,7 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_listed_entry_and_nothing patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), ): never_listed_tool, never_listed_data = await call() - manager._record_listed_tools( + manager.record_listed_tools( petstore, [MCPTool(name="list_pets", description="ADMIN DESC", inputSchema=schema)], ListedToolsCaller(user_api_key_auth=alice), @@ -8156,7 +8156,7 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_guarded_catalog_entry_cl ) manager = mcp_module.global_mcp_server_manager alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice") - manager._record_listed_tools( + manager.record_listed_tools( petstore, [MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=pinned_schema)], ListedToolsCaller(user_api_key_auth=alice), @@ -8206,12 +8206,12 @@ async def test_execute_mcp_tool_hands_openapi_hooks_each_callers_own_listed_entr manager = mcp_module.global_mcp_server_manager guarded = UserAPIKeyAuth(api_key="sk-guarded", user_id="alice") opted_out = UserAPIKeyAuth(api_key="sk-opted-out", user_id="bob") - manager._record_listed_tools( + manager.record_listed_tools( petstore, [MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=schema)], ListedToolsCaller(user_api_key_auth=guarded), ) - manager._record_listed_tools( + manager.record_listed_tools( petstore, [MCPTool(name="getpetbyid", description="Find a SECRET pet", inputSchema=schema)], ListedToolsCaller(user_api_key_auth=opted_out), @@ -8305,7 +8305,7 @@ async def test_execute_mcp_tool_hands_hooks_nothing_for_a_never_listed_operation ) manager = mcp_module.global_mcp_server_manager alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice") - manager._record_listed_tools( + manager.record_listed_tools( petstore, [MCPTool(name="get_pet", description="Fetches pet records. FLAGWORD", inputSchema={"type": "object"})], ListedToolsCaller(user_api_key_auth=alice), From 02f61c9c420b9aa9de10ff673098ad7132b78f5b Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Mon, 5 Oct 2026 03:44:10 -0700 Subject: [PATCH 04/19] chore(cost-map): add azure retirement dates for gpt-4o-realtime-preview-2024-10-01 and jamba-instruct (#44567) Price-Sync: litellm-providers Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> --- litellm/model_prices_and_context_window_backup.json | 8 ++++++++ model_prices_and_context_window.json | 8 ++++++++ 2 files changed, 16 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 2d0985bcb64..e24231b55f5 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -4663,6 +4663,7 @@ "azure/eu/gpt-4o-realtime-preview-2024-10-01": { "cache_creation_input_audio_token_cost": 2.2e-05, "cache_read_input_token_cost": 2.75e-06, + "deprecation_date": "2025-03-26", "input_cost_per_audio_token": 0.00011, "input_cost_per_token": 5.5e-06, "litellm_provider": "azure", @@ -4672,6 +4673,7 @@ "mode": "realtime", "output_cost_per_audio_token": 0.00022, "output_cost_per_token": 2.2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_audio_input": true, "supports_audio_output": true, "supports_function_calling": true, @@ -6306,6 +6308,7 @@ "azure/gpt-4o-realtime-preview-2024-10-01": { "cache_creation_input_audio_token_cost": 2e-05, "cache_read_input_token_cost": 2.5e-06, + "deprecation_date": "2025-03-26", "input_cost_per_audio_token": 0.0001, "input_cost_per_token": 5e-06, "litellm_provider": "azure", @@ -6315,6 +6318,7 @@ "mode": "realtime", "output_cost_per_audio_token": 0.0002, "output_cost_per_token": 2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_audio_input": true, "supports_audio_output": true, "supports_function_calling": true, @@ -11093,6 +11097,7 @@ "azure/us/gpt-4o-realtime-preview-2024-10-01": { "cache_creation_input_audio_token_cost": 2.2e-05, "cache_read_input_token_cost": 2.75e-06, + "deprecation_date": "2025-03-26", "input_cost_per_audio_token": 0.00011, "input_cost_per_token": 5.5e-06, "litellm_provider": "azure", @@ -11102,6 +11107,7 @@ "mode": "realtime", "output_cost_per_audio_token": 0.00022, "output_cost_per_token": 2.2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_audio_input": true, "supports_audio_output": true, "supports_function_calling": true, @@ -12556,6 +12562,7 @@ "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models" }, "azure_ai/jamba-instruct": { + "deprecation_date": "2025-03-01", "input_cost_per_token": 5e-07, "litellm_provider": "azure_ai", "max_input_tokens": 70000, @@ -12563,6 +12570,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 7e-07, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true }, "azure_ai/kimi-k2.5": { diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 2d0985bcb64..e24231b55f5 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -4663,6 +4663,7 @@ "azure/eu/gpt-4o-realtime-preview-2024-10-01": { "cache_creation_input_audio_token_cost": 2.2e-05, "cache_read_input_token_cost": 2.75e-06, + "deprecation_date": "2025-03-26", "input_cost_per_audio_token": 0.00011, "input_cost_per_token": 5.5e-06, "litellm_provider": "azure", @@ -4672,6 +4673,7 @@ "mode": "realtime", "output_cost_per_audio_token": 0.00022, "output_cost_per_token": 2.2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_audio_input": true, "supports_audio_output": true, "supports_function_calling": true, @@ -6306,6 +6308,7 @@ "azure/gpt-4o-realtime-preview-2024-10-01": { "cache_creation_input_audio_token_cost": 2e-05, "cache_read_input_token_cost": 2.5e-06, + "deprecation_date": "2025-03-26", "input_cost_per_audio_token": 0.0001, "input_cost_per_token": 5e-06, "litellm_provider": "azure", @@ -6315,6 +6318,7 @@ "mode": "realtime", "output_cost_per_audio_token": 0.0002, "output_cost_per_token": 2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_audio_input": true, "supports_audio_output": true, "supports_function_calling": true, @@ -11093,6 +11097,7 @@ "azure/us/gpt-4o-realtime-preview-2024-10-01": { "cache_creation_input_audio_token_cost": 2.2e-05, "cache_read_input_token_cost": 2.75e-06, + "deprecation_date": "2025-03-26", "input_cost_per_audio_token": 0.00011, "input_cost_per_token": 5.5e-06, "litellm_provider": "azure", @@ -11102,6 +11107,7 @@ "mode": "realtime", "output_cost_per_audio_token": 0.00022, "output_cost_per_token": 2.2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_audio_input": true, "supports_audio_output": true, "supports_function_calling": true, @@ -12556,6 +12562,7 @@ "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models" }, "azure_ai/jamba-instruct": { + "deprecation_date": "2025-03-01", "input_cost_per_token": 5e-07, "litellm_provider": "azure_ai", "max_input_tokens": 70000, @@ -12563,6 +12570,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 7e-07, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true }, "azure_ai/kimi-k2.5": { From 9cc15e9320076d51ae1bf53dff5132bfe0194fd8 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 5 Oct 2026 15:03:37 +0000 Subject: [PATCH 05/19] feat(ui): add persistent columns and loading skeletons to Lens runs (#44579) Co-authored-by: Yujong Lee Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../traces/list/AgentTracesTable.test.tsx | 97 ++++++++++++++++++- .../lens/traces/list/AgentTracesTable.tsx | 68 +++++++++---- .../shared/DataTable/DataTableViewOptions.tsx | 33 ++++++- .../components/shared/InspectorTable.test.tsx | 54 +++++++++++ .../src/components/shared/InspectorTable.tsx | 34 ++++++- 5 files changed, 257 insertions(+), 29 deletions(-) diff --git a/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.test.tsx b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.test.tsx index dff968334f7..a99c44a98a6 100644 --- a/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.test.tsx @@ -1,5 +1,6 @@ -import { fireEvent, render, screen } from "@testing-library/react"; -import { describe, expect, it, vi } from "vitest"; +import { fireEvent, render, screen, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; import { renderWithProviders } from "../../../../../tests/test-utils"; import { Inspector } from "@/components/shared/Inspector"; @@ -51,6 +52,30 @@ describe("AgentTracesTable empty state", () => { }); }); +describe("AgentTracesTable loading state", () => { + it("fills the first page load with skeleton rows instead of an empty table", () => { + render( + inList( + , + ), + ); + expect(screen.getByRole("status")).toHaveTextContent("Loading runs…"); + const placeholders = screen.getAllByTestId("runs-placeholder"); + expect(placeholders.length).toBeGreaterThanOrEqual(8); + const columnCount = screen.getAllByRole("columnheader").length; + expect(within(placeholders[0]).getAllByRole("cell", { hidden: true })).toHaveLength(columnCount); + expect(screen.queryByText(/No runs/)).not.toBeInTheDocument(); + }); +}); + describe("AgentTracesTable virtualization", () => { const template = (traceList as TracePage).data[0] as TraceSummary; const manyRuns: TraceSummary[] = Array.from({ length: 500 }, (_, i) => ({ @@ -84,3 +109,71 @@ describe("AgentTracesTable virtualization", () => { expect(screen.queryByText("question 0")).not.toBeInTheDocument(); }); }); + +describe("AgentTracesTable column picker", () => { + const runs = (traceList as TracePage).data as TraceSummary[]; + const renderRuns = () => + renderWithProviders( + inList( + , + ), + ); + const headers = () => screen.getAllByRole("columnheader").map((header) => header.textContent); + + beforeEach(() => localStorage.clear()); + + it("hides a column from the header picker and keeps it hidden after a remount", async () => { + const user = userEvent.setup(); + const { unmount } = renderRuns(); + const before = headers().length; + expect(headers()).toContain("Cost"); + + await user.click(screen.getByRole("button", { name: "Columns" })); + expect(screen.queryByTestId("view-option-time")).not.toBeInTheDocument(); + expect(screen.queryByTestId("view-option-agent")).not.toBeInTheDocument(); + expect(screen.queryByTestId("view-option-open")).not.toBeInTheDocument(); + await user.click(await screen.findByTestId("view-option-cost")); + + expect(headers()).not.toContain("Cost"); + expect(headers()).toHaveLength(before - 1); + const [firstRow] = screen.getAllByTestId("agent-trace-row"); + expect(within(firstRow).getAllByRole("cell")).toHaveLength(before - 1); + + unmount(); + renderRuns(); + expect(headers()).not.toContain("Cost"); + expect(headers()).toHaveLength(before - 1); + }); + + it("shapes loading skeletons to the columns still visible", async () => { + const user = userEvent.setup(); + const { unmount } = renderRuns(); + await user.click(screen.getByRole("button", { name: "Columns" })); + await user.click(await screen.findByTestId("view-option-cost")); + unmount(); + + render( + inList( + , + ), + ); + const columnCount = screen.getAllByRole("columnheader").length; + const [placeholder] = screen.getAllByTestId("runs-placeholder"); + expect(within(placeholder).getAllByRole("cell", { hidden: true })).toHaveLength(columnCount); + }); +}); diff --git a/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.tsx b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.tsx index 7d860837bab..1d1334b7881 100644 --- a/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesTable.tsx @@ -5,8 +5,11 @@ import { ArrowDown, ChevronRight } from "lucide-react"; import { useEffect } from "react"; import { useInView } from "react-intersection-observer"; +import { DataTableViewOptions } from "@/components/shared/DataTable/DataTableViewOptions"; +import { usePersistedColumnVisibility } from "@/components/shared/DataTable/usePersistedColumnVisibility"; import { InspectorTable, useInspectorTable } from "@/components/shared/InspectorTable"; import { Button } from "@/components/ui/button"; +import { Skeleton } from "@/components/ui/skeleton"; import { formatActivityTimestamp, formatRunTimestamp, localTimeZoneAbbreviation } from "@/utils/activityTimestamp"; import { SpanIcon } from "../ui/SpanIcon"; @@ -39,6 +42,7 @@ const runKey = (run: TraceSummary): string => run.trace_ref || run.trace_id; const PREFETCH_MARGIN = "0px 0px 480px 0px"; const PLACEHOLDER_ROWS = [0, 1, 2]; +const SKELETON_ROWS = Array.from({ length: 12 }, (_, i) => i); const ROW_HEIGHT = 36; const MUTED_NUM = "font-mono text-muted-foreground"; const NUM = "font-mono text-foreground"; @@ -79,6 +83,7 @@ const RUN_COLUMNS: ColumnDef[] = [ { id: "time", size: 170, + enableHiding: false, header: () => ( Time @@ -92,9 +97,27 @@ const RUN_COLUMNS: ColumnDef[] = [ {formatRunTimestamp(row.original.start_time)} ), - meta: { className: "font-mono tabular-nums text-muted-foreground" }, + meta: { + title: "Time", + className: "font-mono tabular-nums text-muted-foreground", + renderSkeleton: () => , + }, + }, + { + id: "agent", + size: 160, + enableHiding: false, + header: "Agent", + cell: ({ row }) => , + meta: { + renderSkeleton: () => ( +
+ + +
+ ), + }, }, - { id: "agent", size: 160, header: "Agent", cell: ({ row }) => }, { id: "input", header: "Input", cell: ({ row }) => }, { id: "agents", @@ -139,26 +162,15 @@ const RUN_COLUMNS: ColumnDef[] = [ { id: "open", size: 32, - header: "", + enableHiding: false, + header: ({ table }) => , cell: () => , - meta: { className: "px-0" }, + meta: { className: "px-0", headerClassName: "px-1", renderSkeleton: () => null }, }, ]; -function PlaceholderRow({ rowRef }: { rowRef?: (node: Element | null) => void }) { - return ( - - -
- - -
- - -
- - - ); +function PlaceholderRow({ index, rowRef }: { index: number; rowRef?: (node: Element | null) => void }) { + return ; } function LoadMoreRows({ isFetching, onLoadMore }: { isFetching: boolean; onLoadMore: () => void }) { @@ -167,7 +179,9 @@ function LoadMoreRows({ isFetching, onLoadMore }: { isFetching: boolean; onLoadM useEffect(() => { if (nearTail && !isFetching) onLoadMore(); }, [nearTail, isFetching, onLoadMore]); - return PLACEHOLDER_ROWS.map((row) => ); + return PLACEHOLDER_ROWS.map((row) => ( + + )); } function EmptyRuns({ rangeEmpty, onSetUpTracing }: { rangeEmpty: boolean; onSetUpTracing: () => void }) { @@ -207,11 +221,14 @@ export function AgentTracesTable({ const isEmpty = settled && !hasMore && traces.length === 0; const canContinue = settled && hasMore; const autoContinue = canContinue && traces.length > 0; + const { columnVisibility, onColumnVisibilityChange } = usePersistedColumnVisibility("lens-traces"); const tableOptions: TableOptions = { data: traces, columns: RUN_COLUMNS, getRowId: runKey, autoResetAll: false, + state: { columnVisibility }, + onColumnVisibilityChange, getCoreRowModel: getCoreRowModel(), }; const table = useReactTable(tableOptions); @@ -221,7 +238,12 @@ export function AgentTracesTable({ rowHeight={() => ROW_HEIGHT} - after={autoContinue && } + after={ + <> + {isLoading && SKELETON_ROWS.map((row) => )} + {autoContinue && } + + } > {(row) => ( - {isLoading &&
Loading runs…
} + {isLoading && ( +

+ Loading runs… +

+ )} {error && (
diff --git a/ui/litellm-dashboard/src/components/shared/DataTable/DataTableViewOptions.tsx b/ui/litellm-dashboard/src/components/shared/DataTable/DataTableViewOptions.tsx index ca60e42d9aa..86d99fd72a6 100644 --- a/ui/litellm-dashboard/src/components/shared/DataTable/DataTableViewOptions.tsx +++ b/ui/litellm-dashboard/src/components/shared/DataTable/DataTableViewOptions.tsx @@ -5,14 +5,21 @@ import type { Table } from "@tanstack/react-table"; import { Check, Columns3 } from "lucide-react"; import { Button } from "@/components/ui/button"; +import { cn } from "@/lib/cva.config"; interface DataTableViewOptionsProps { table: Table; label?: string; + iconOnly?: boolean; className?: string; } -export function DataTableViewOptions({ table, label = "View", className }: DataTableViewOptionsProps) { +export function DataTableViewOptions({ + table, + label = "View", + iconOnly = false, + className, +}: DataTableViewOptionsProps) { const hideableColumns = table.getAllLeafColumns().filter((column) => column.getCanHide()); if (hideableColumns.length === 0) { @@ -23,10 +30,26 @@ export function DataTableViewOptions({ table, label = "View", className } - - {label} - + iconOnly ? ( + + ) : ( + + ) } /> diff --git a/ui/litellm-dashboard/src/components/shared/InspectorTable.test.tsx b/ui/litellm-dashboard/src/components/shared/InspectorTable.test.tsx index 416d31bc940..8201248c078 100644 --- a/ui/litellm-dashboard/src/components/shared/InspectorTable.test.tsx +++ b/ui/litellm-dashboard/src/components/shared/InspectorTable.test.tsx @@ -104,3 +104,57 @@ describe("InspectorTable tree", () => { expect(screen.getByText("Open: b1")).toBeInTheDocument(); }); }); + +const SKELETON_COLUMNS: ColumnDef[] = [ + { id: "name", header: "Name", meta: { renderSkeleton: () => } }, + { id: "size", header: "Size", meta: { numeric: true } }, +]; + +function SkeletonTree({ visibility }: { visibility: Record }) { + const tableOptions: TableOptions = { + data: TREE, + columns: SKELETON_COLUMNS, + getRowId: (node) => node.id, + state: { columnVisibility: visibility }, + getCoreRowModel: getCoreRowModel(), + }; + const table = useReactTable(tableOptions); + return ( + node.id} + selected={null} + onSelectedChange={() => {}} + noun="node" + storageKey="inspector-skeleton-test" + > + + + + + rowHeight={() => 36} + after={} + > + {(row) => } + + + + + ); +} + +describe("InspectorTable.SkeletonRow", () => { + it("renders one hidden cell per visible leaf column and honors renderSkeleton", () => { + render(); + const [row] = screen.getAllByTestId("skeleton-row"); + expect(row).toHaveAttribute("aria-hidden", "true"); + expect(within(row).getAllByRole("cell", { hidden: true })).toHaveLength(2); + expect(screen.getAllByTestId("custom-skeleton")).toHaveLength(1); + }); + + it("follows column visibility", () => { + render(); + const [row] = screen.getAllByTestId("skeleton-row"); + expect(within(row).getAllByRole("cell", { hidden: true })).toHaveLength(1); + }); +}); diff --git a/ui/litellm-dashboard/src/components/shared/InspectorTable.tsx b/ui/litellm-dashboard/src/components/shared/InspectorTable.tsx index d5443fe9ac1..059ad1673ee 100644 --- a/ui/litellm-dashboard/src/components/shared/InspectorTable.tsx +++ b/ui/litellm-dashboard/src/components/shared/InspectorTable.tsx @@ -8,6 +8,7 @@ import { ChevronRight } from "lucide-react"; import { createContext, Fragment, useContext, useState, type ComponentProps, type ReactNode } from "react"; import { Inspector } from "@/components/shared/Inspector"; +import { Skeleton } from "@/components/ui/skeleton"; import { TableBody, TableCell, TableHead, TableHeader } from "@/components/ui/table"; import { cn } from "@/lib/cva.config"; @@ -137,6 +138,37 @@ function Row({ row, item, className, ...props }: RowProps) { ); } +const SKELETON_WIDTHS = ["w-[58%]", "w-[44%]", "w-[70%]", "w-[50%]", "w-[64%]", "w-[48%]"] as const; + +type SkeletonRowProps = ComponentProps<"tr"> & { readonly index: number }; + +/** One placeholder row shaped by the visible columns: `meta.renderSkeleton` wins, numeric cells right-align. */ +function SkeletonRow({ index, className, ...props }: SkeletonRowProps) { + const { table } = useInspectorTable(); + return ( + + {table.getVisibleLeafColumns().map((column, position) => { + const meta = column.columnDef.meta; + return ( + + {meta?.renderSkeleton ? ( + meta.renderSkeleton() + ) : ( + + )} + + ); + })} + + ); +} + interface IndentProps { readonly row: TanStackRow; readonly toggleLabel?: (expanded: boolean) => string; @@ -176,4 +208,4 @@ function Indent({ row, toggleLabel = (expanded) => (expanded ? "Collapse" : " ); } -export const InspectorTable = { Root, Grid, Header, Body, Row, Indent } as const; +export const InspectorTable = { Root, Grid, Header, Body, Row, SkeletonRow, Indent } as const; From adb59a7af8a709801f19056759a3a55e93d750a1 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 5 Oct 2026 08:42:19 -0700 Subject: [PATCH 06/19] fix(ui): remove Top models by task card from Model Leaderboard (#44502) Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- .../_components/ModelInsightsView.test.tsx | 57 +-------- .../_components/ModelInsightsView.tsx | 118 +----------------- .../_components/modelInsightsData.ts | 10 -- 3 files changed, 8 insertions(+), 177 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.test.tsx index 37d265c4e35..7aad92f9814 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.test.tsx @@ -20,7 +20,6 @@ vi.mock("recharts", () => ({
), CartesianGrid: () => null, - Treemap: () => null, XAxis: () => null, YAxis: () => null, })); @@ -45,48 +44,26 @@ const response = { daily_totals: [{ date: "2026-09-28", spend: 2.5, prompt_tokens: 1000, completion_tokens: 2000, requests: 12 }], }; -const taskResponse = { - start_date: "2025-09-29", - end_date: "2026-09-28", - tasks: [ - { - task_type: "code_generation", - label: "Code Generation", - category: "Code", - value: 2.5, - share: 100, - leader: "fast-chat", - provider: "openai", - }, - ], -}; - -const mockApi = (tasks: unknown = taskResponse) => - vi - .mocked(apiClient.get) - .mockImplementation((path: string) => - path === "/model-insights/tasks" ? (tasks as Promise) : Promise.resolve(response), - ); - describe("ModelInsightsView", () => { beforeEach(() => { vi.mocked(apiClient.get).mockReset(); - mockApi(Promise.resolve(taskResponse)); + vi.mocked(apiClient.get).mockResolvedValue(response); }); - it("shows the ranking with share and the task legend from the API response", async () => { + it("shows the ranking with share from the API response", async () => { render(); expect(await screen.findByText("fast-chat")).toBeInTheDocument(); expect(screen.getByText("by openai")).toBeInTheDocument(); - expect(await screen.findByText("Code")).toBeInTheDocument(); - expect(screen.getAllByText("100.0%")).toHaveLength(2); + expect(screen.getAllByText("100.0%")).toHaveLength(1); + expect(screen.queryByText("Top models by task")).not.toBeInTheDocument(); expect(screen.getByRole("tab", { name: "tokens" })).toHaveAttribute("aria-selected", "true"); expect(screen.getByRole("tab", { name: "log" })).toBeInTheDocument(); expect(apiClient.get).toHaveBeenCalledWith("/model-insights", { accessToken: "token", query: { metric: "tokens" }, }); + expect(apiClient.get).not.toHaveBeenCalledWith("/model-insights/tasks", expect.anything()); }); it("refetches with the selected metric so top models are ranked by it", async () => { @@ -103,24 +80,6 @@ describe("ModelInsightsView", () => { ); }); - it("does not refetch the task breakdown when the chart metric changes", async () => { - render(); - await screen.findByText("Code"); - const taskCalls = () => - vi.mocked(apiClient.get).mock.calls.filter(([path]) => path === "/model-insights/tasks").length; - const before = taskCalls(); - - await userEvent.click(screen.getByRole("tab", { name: "requests" })); - await waitFor(() => - expect(apiClient.get).toHaveBeenCalledWith("/model-insights", { - accessToken: "token", - query: { metric: "requests" }, - }), - ); - - expect(taskCalls()).toBe(before); - }); - it("shows the API error instead of loading forever", async () => { vi.mocked(apiClient.get).mockRejectedValue(new Error("Only proxy admins can view deployment-wide model insights")); render(); @@ -133,11 +92,7 @@ describe("ModelInsightsView", () => { render(); await screen.findByText("fast-chat"); let resolve: (value: typeof response) => void = () => {}; - vi.mocked(apiClient.get).mockImplementation((path: string) => - path === "/model-insights/tasks" - ? Promise.resolve(taskResponse) - : new Promise((done) => (resolve = done as typeof resolve)), - ); + vi.mocked(apiClient.get).mockImplementation(() => new Promise((done) => (resolve = done as typeof resolve))); await userEvent.click(screen.getByRole("tab", { name: "spend" })); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.tsx index ab406597589..21a707b8c42 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/ModelInsightsView.tsx @@ -2,8 +2,8 @@ import { Page } from "@/components/shared/Page"; import React from "react"; -import { Bar, BarChart, CartesianGrid, Treemap, XAxis, YAxis } from "recharts"; -import { ArrowDownRight, ArrowUpRight, BarChart3, Layers, Minus } from "lucide-react"; +import { Bar, BarChart, CartesianGrid, XAxis, YAxis } from "recharts"; +import { ArrowDownRight, ArrowUpRight, BarChart3, Minus } from "lucide-react"; import { apiClient } from "@/components/networking"; import { extractErrorMessage } from "@/utils/errorUtils"; @@ -12,7 +12,6 @@ import { PageHeader, PageHeaderDescription, PageHeaderTitle } from "@/components import { Alert, AlertDescription, AlertTitle } from "@/components/ui/alert"; import { Card, CardContent, CardDescription, CardHeader, CardTitle } from "@/components/ui/card"; import { ChartConfig, ChartContainer, ChartTooltip, ChartTooltipContent } from "@/components/ui/chart"; -import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; import { Skeleton } from "@/components/ui/skeleton"; import { Tabs, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { @@ -22,8 +21,6 @@ import { Granularity, Metric, ModelInsightsResponse, - ModelInsightTasksResponse, - TaskSummary, modelOrder, rankModels, RankedModel, @@ -41,13 +38,6 @@ const PALETTE = [ "#6366f1", "#f97316", ]; -const FALLBACK_COLOR = "#64748b"; -const CATEGORY_COLORS: Record = { - General: "#ee8650", - Agent: "#7666e4", - Code: "#5fb074", - Data: "#3b82f6", -}; const SCALES = ["linear", "log"] as const; const GRANULARITIES = ["day", "week"] as const; const GRANULARITY_LABELS: Record = { day: "Daily", week: "Weekly" }; @@ -90,37 +80,11 @@ const RankingRow = ({ model, rank }: { model: RankedModel; rank: number }) => ( ); -type TileProps = TaskSummary & { x: number; y: number; width: number; height: number; index: number }; - -const TaskTileContent = ({ x, y, width, height, category, label, leader }: TileProps) => { - if (width <= 0 || height <= 0) return null; - const color = CATEGORY_COLORS[category] ?? FALLBACK_COLOR; - const fits = width > 90 && height > 44; - return ( - - - {fits && ( - <> - - {label} - - - {leader} - - - )} - - ); -}; - export default function ModelInsightsView({ accessToken }: { accessToken: string | null }) { const [loaded, setLoaded] = React.useState<{ metric: Metric; response: ModelInsightsResponse } | null>(null); const [metric, setMetric] = React.useState("tokens"); const [scale, setScale] = React.useState("linear"); const [granularity, setGranularity] = React.useState("day"); - const [taskMetric, setTaskMetric] = React.useState("spend"); - const [taskData, setTaskData] = React.useState(null); - const [taskError, setTaskError] = React.useState(null); const [error, setError] = React.useState(null); React.useEffect(() => { @@ -141,24 +105,6 @@ export default function ModelInsightsView({ accessToken }: { accessToken: string }; }, [accessToken, metric]); - React.useEffect(() => { - if (!accessToken) return; - let cancelled = false; - apiClient - .get("/model-insights/tasks", { accessToken, query: { metric: taskMetric } }) - .then((response) => { - if (cancelled) return; - setTaskError(null); - setTaskData(response); - }) - .catch((err: unknown) => { - if (!cancelled) setTaskError(extractErrorMessage(err)); - }); - return () => { - cancelled = true; - }; - }, [accessToken, taskMetric]); - const data = loaded?.response ?? null; const shown = loaded?.metric ?? metric; const isStale = loaded !== null && loaded.metric !== metric; @@ -176,15 +122,6 @@ export default function ModelInsightsView({ accessToken }: { accessToken: string () => (data ? rankModels(data.top_models, data.daily, shown, range) : []), [data, shown, range], ); - const tiles = React.useMemo(() => taskData?.tasks ?? [], [taskData]); - const categoryShares = React.useMemo( - () => - [...new Set(tiles.map((tile) => tile.category))].map((category) => ({ - category, - share: tiles.filter((tile) => tile.category === category).reduce((sum, tile) => sum + tile.share, 0), - })), - [tiles], - ); if (error) { return ( @@ -317,57 +254,6 @@ export default function ModelInsightsView({ accessToken }: { accessToken: string - - -
- - Top models by task - - - Each task's share of {METRIC_LABELS[taskMetric]}, labelled with its leading model - -
- -
- - {taskError && ( - - Could not load tasks - {taskError} - - )} - - ({ ...tile, name: tile.task_type }))} - dataKey="value" - isAnimationActive={false} - content={} - /> - -
    - {categoryShares.map(({ category, share }) => ( -
  • - - {category} - {share.toFixed(1)}% -
  • - ))} -
-
-
- Cost per session diff --git a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/modelInsightsData.ts b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/modelInsightsData.ts index 0e1a5fa7aac..9e0aecc9e8e 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/modelInsightsData.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/model-insights/_components/modelInsightsData.ts @@ -21,16 +21,6 @@ export type ModelInsightsResponse = { daily_totals: DailyTotal[]; top_models: ModelMetric[]; }; -export type TaskSummary = { - task_type: string; - label: string; - category: string; - value: number; - share: number; - leader: string; - provider: string; -}; -export type ModelInsightTasksResponse = { start_date: string; end_date: string; tasks: TaskSummary[] }; export type RankedModel = { model_group: string; provider: string; share: number; delta: number }; export type Granularity = "day" | "week"; From fdc0e7dc2b98d68daa5b83e61b7e2ff9e7c348c1 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 5 Oct 2026 16:04:44 +0000 Subject: [PATCH 07/19] refactor(rust): track ClickHouse migrations in a checksummed ledger (#44580) * refactor(rust): track ClickHouse migrations in a checksummed ledger Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(clickhouse): harden migration execution and reuse retention SQL * refactor(rust): track ClickHouse migrations in a checksummed ledger Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * docs(rust): keep ClickHouse retention TTLs in the current policy list Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Yujong Lee Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm-rust/Cargo.lock | 26 +- litellm-rust/Cargo.toml | 2 - litellm-rust/crates/migrate-macros/Cargo.toml | 19 - .../crates/migrate-macros/src/error.rs | 21 - litellm-rust/crates/migrate-macros/src/lib.rs | 199 -------- litellm-rust/crates/migrate/Cargo.toml | 12 - litellm-rust/crates/migrate/README.md | 5 - litellm-rust/crates/migrate/src/lib.rs | 8 - .../tests/fixtures/migrations/10_tenth.sql | 1 - .../tests/fixtures/migrations/1_first.sql | 1 - .../tests/fixtures/migrations/2_second.sql | 1 - litellm-rust/crates/migrate/tests/migrate.rs | 21 - litellm-rust/crates/python-bridge/Cargo.toml | 1 + .../crates/python-bridge/src/routes/traces.rs | 12 +- .../crates/storage-clickhouse/AGENTS.md | 4 + .../crates/storage-clickhouse/Cargo.toml | 3 + .../crates/storage-clickhouse/src/error.rs | 2 +- .../crates/storage-clickhouse/src/lib.rs | 2 + .../crates/storage-clickhouse/src/migrate.rs | 308 ++++++++++++ .../storage-clickhouse/tests/migrations.rs | 476 ++++++++++++++++++ .../crates/traces-clickhouse/AGENTS.md | 3 +- .../crates/traces-clickhouse/Cargo.toml | 2 +- .../crates/traces-clickhouse/src/error.rs | 6 +- .../crates/traces-clickhouse/src/lib.rs | 3 +- .../crates/traces-clickhouse/src/schema.rs | 148 +++++- .../traces-clickhouse/tests/migrations.rs | 226 ++++++++- tests/test_litellm_rust/test_traces.py | 9 +- 27 files changed, 1173 insertions(+), 348 deletions(-) delete mode 100644 litellm-rust/crates/migrate-macros/Cargo.toml delete mode 100644 litellm-rust/crates/migrate-macros/src/error.rs delete mode 100644 litellm-rust/crates/migrate-macros/src/lib.rs delete mode 100644 litellm-rust/crates/migrate/Cargo.toml delete mode 100644 litellm-rust/crates/migrate/README.md delete mode 100644 litellm-rust/crates/migrate/src/lib.rs delete mode 100644 litellm-rust/crates/migrate/tests/fixtures/migrations/10_tenth.sql delete mode 100644 litellm-rust/crates/migrate/tests/fixtures/migrations/1_first.sql delete mode 100644 litellm-rust/crates/migrate/tests/fixtures/migrations/2_second.sql delete mode 100644 litellm-rust/crates/migrate/tests/migrate.rs create mode 100644 litellm-rust/crates/storage-clickhouse/src/migrate.rs create mode 100644 litellm-rust/crates/storage-clickhouse/tests/migrations.rs diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index f7b667c8ab2..934a1f0bb6b 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -4085,26 +4085,6 @@ dependencies = [ "strum", ] -[[package]] -name = "litellm-migrate" -version = "0.1.0" -dependencies = [ - "litellm-migrate-macros", - "rstest", -] - -[[package]] -name = "litellm-migrate-macros" -version = "0.1.0" -dependencies = [ - "proc-macro2", - "quote", - "rstest", - "syn 2.0.119", - "tempfile", - "thiserror 2.0.19", -] - [[package]] name = "litellm-model-catalog" version = "0.1.0" @@ -4170,6 +4150,7 @@ dependencies = [ "serde_json", "serde_with", "sha2 0.10.9", + "sqlx", "strum", "thiserror 2.0.19", "tokio", @@ -4368,6 +4349,9 @@ dependencies = [ "rstest", "serde", "serde_json", + "serde_with", + "sqlx", + "testcontainers-modules", "thiserror 2.0.19", "tokio", "url", @@ -4498,7 +4482,6 @@ dependencies = [ "hmac 0.12.1", "jsonschema", "litellm-http", - "litellm-migrate", "litellm-storage-clickhouse", "litellm-traces", "litellm-traces-cache", @@ -4509,6 +4492,7 @@ dependencies = [ "serde", "serde_json", "sha2 0.10.9", + "sqlx", "strum", "testcontainers-modules", "thiserror 2.0.19", diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index b0766f11e87..6ef90c59d2f 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -16,8 +16,6 @@ litellm-traces = { path = "crates/traces" } litellm-traces-cache = { path = "crates/traces-cache" } litellm-traces-clickhouse = { path = "crates/traces-clickhouse" } litellm-storage-clickhouse = { path = "crates/storage-clickhouse" } -litellm-migrate = { path = "crates/migrate" } -litellm-migrate-macros = { path = "crates/migrate-macros" } litellm-core = { path = "crates/core" } litellm-gateway-mcp = { path = "crates/gateway-mcp" } litellm-gateway = { path = "crates/gateway" } diff --git a/litellm-rust/crates/migrate-macros/Cargo.toml b/litellm-rust/crates/migrate-macros/Cargo.toml deleted file mode 100644 index 5cd68415ca2..00000000000 --- a/litellm-rust/crates/migrate-macros/Cargo.toml +++ /dev/null @@ -1,19 +0,0 @@ -[package] -name = "litellm-migrate-macros" -version = "0.1.0" -edition.workspace = true -license.workspace = true -repository.workspace = true - -[lib] -proc-macro = true - -[dependencies] -proc-macro2.workspace = true -quote.workspace = true -syn = { workspace = true, features = ["parsing", "printing", "proc-macro"] } -thiserror.workspace = true - -[dev-dependencies] -rstest.workspace = true -tempfile.workspace = true diff --git a/litellm-rust/crates/migrate-macros/src/error.rs b/litellm-rust/crates/migrate-macros/src/error.rs deleted file mode 100644 index 9833009517b..00000000000 --- a/litellm-rust/crates/migrate-macros/src/error.rs +++ /dev/null @@ -1,21 +0,0 @@ -use std::io; - -#[derive(Debug, thiserror::Error)] -pub enum Error { - #[error("could not read migrations directory `{path}`")] - ReadDirectory { - path: String, - #[source] - source: io::Error, - }, - #[error( - "migration name `{name}` must be `_.sql` with a `[a-z0-9_]` description" - )] - InvalidName { name: String }, - #[error("migration version `{version}` is declared more than once")] - DuplicateVersion { version: u64 }, - #[error("migrations directory `{path}` contains no migrations")] - Empty { path: String }, - #[error("migration path `{path}` is not valid UTF-8")] - NonUtf8Path { path: String }, -} diff --git a/litellm-rust/crates/migrate-macros/src/lib.rs b/litellm-rust/crates/migrate-macros/src/lib.rs deleted file mode 100644 index 501f59e6fc2..00000000000 --- a/litellm-rust/crates/migrate-macros/src/lib.rs +++ /dev/null @@ -1,199 +0,0 @@ -mod error; - -use std::path::{Path, PathBuf}; - -use error::Error; -use proc_macro::TokenStream; -use quote::quote; -use syn::LitStr; - -struct Entry { - version: u64, - description: String, - path: PathBuf, -} - -fn resolve(dir: &Path) -> Result, Error> { - let mut entries = Vec::new(); - let files = std::fs::read_dir(dir).map_err(|source| Error::ReadDirectory { - path: dir.display().to_string(), - source, - })?; - for file in files { - let file = file.map_err(|source| Error::ReadDirectory { - path: dir.display().to_string(), - source, - })?; - let path = file.path(); - let name = path - .file_name() - .and_then(|name| name.to_str()) - .ok_or_else(|| Error::NonUtf8Path { - path: path.display().to_string(), - })? - .to_owned(); - let invalid = || Error::InvalidName { name: name.clone() }; - let stem = name - .strip_suffix(".sql") - .filter(|_| file.file_type().is_ok_and(|kind| kind.is_file())) - .and_then(|stem| stem.split_once('_')) - .filter(|(version, description)| { - !version.is_empty() - && version.bytes().all(|b| b.is_ascii_digit()) - && !description.is_empty() - && description - .bytes() - .all(|b| b.is_ascii_lowercase() || b.is_ascii_digit() || b == b'_') - }) - .ok_or_else(invalid)?; - let version = stem.0.parse::().map_err(|_| invalid())?; - entries.push(Entry { - version, - description: stem.1.to_owned(), - path, - }); - } - if entries.is_empty() { - return Err(Error::Empty { - path: dir.display().to_string(), - }); - } - entries.sort_by_key(|entry| entry.version); - for pair in entries.windows(2) { - if pair[0].version == pair[1].version { - return Err(Error::DuplicateVersion { - version: pair[0].version, - }); - } - } - Ok(entries) -} - -fn resolve_input(lit: &LitStr) -> Result, Error> { - let root = std::env::var("CARGO_MANIFEST_DIR") - .map(PathBuf::from) - .unwrap_or_default(); - let dir = root.join(lit.value()); - let dir = dir.canonicalize().map_err(|source| Error::ReadDirectory { - path: dir.display().to_string(), - source, - })?; - if dir.to_str().is_none() { - return Err(Error::NonUtf8Path { - path: dir.display().to_string(), - }); - } - resolve(&dir) -} - -#[proc_macro] -pub fn migrate(input: TokenStream) -> TokenStream { - let lit = syn::parse_macro_input!(input as LitStr); - match resolve_input(&lit) { - Ok(entries) => { - let migrations = entries.iter().map(|entry| { - let version = entry.version; - let description = &entry.description; - let path = entry - .path - .to_str() - .expect("canonical migration path is UTF-8"); - quote! { - ::litellm_migrate::Migration { - version: #version, - description: #description, - sql: ::core::include_str!(#path), - } - } - }); - quote! { &[#(#migrations),*] }.into() - } - Err(err) => syn::Error::new(lit.span(), err).to_compile_error().into(), - } -} - -#[cfg(test)] -mod tests { - use std::fs; - - use rstest::rstest; - use tempfile::TempDir; - - use super::{Error, resolve}; - - fn migrations_dir(files: &[&str]) -> TempDir { - let dir = TempDir::new().expect("tempdir"); - for file in files { - fs::write(dir.path().join(file), "SELECT 1").expect("write fixture"); - } - dir - } - - #[rstest] - fn orders_versions_numerically() { - let dir = migrations_dir(&["10_tenth.sql", "2_second.sql", "1_first.sql"]); - let entries = resolve(dir.path()).expect("resolves"); - let versions: Vec = entries.iter().map(|entry| entry.version).collect(); - let descriptions: Vec<&str> = entries - .iter() - .map(|entry| entry.description.as_str()) - .collect(); - assert_eq!(versions, [1, 2, 10]); - assert_eq!(descriptions, ["first", "second", "tenth"]); - } - - #[rstest] - #[case::dash_in_version(&["0001-dash.sql"])] - #[case::not_sql(&["notes.txt"])] - #[case::empty_description(&["0001_.sql"])] - #[case::non_digit_version(&["x_name.sql"])] - #[case::uppercase_description(&["0001_Upper.sql"])] - #[case::no_underscore(&["0001.sql"])] - #[case::plus_sign_version(&["+10_add.sql"])] - fn rejects_invalid_names(#[case] files: &[&str]) { - let dir = migrations_dir(files); - assert!(matches!( - resolve(dir.path()), - Err(Error::InvalidName { .. }) - )); - } - - #[rstest] - fn rejects_subdirectories() { - let dir = migrations_dir(&["0001_a.sql"]); - fs::create_dir(dir.path().join("0002_b.sql")).expect("subdir"); - assert!(matches!( - resolve(dir.path()), - Err(Error::InvalidName { .. }) - )); - } - - #[cfg(unix)] - #[rstest] - fn rejects_symlinks() { - let dir = migrations_dir(&["0001_a.sql"]); - let target = TempDir::new().expect("tempdir"); - let target_file = target.path().join("real.sql"); - fs::write(&target_file, "SELECT 2").expect("write fixture"); - std::os::unix::fs::symlink(&target_file, dir.path().join("0002_b.sql")).expect("symlink"); - assert!(matches!( - resolve(dir.path()), - Err(Error::InvalidName { .. }) - )); - } - - #[rstest] - fn rejects_duplicate_versions() { - let dir = migrations_dir(&["0001_a.sql", "1_b.sql"]); - assert!(matches!( - resolve(dir.path()), - Err(Error::DuplicateVersion { version: 1 }) - )); - } - - #[rstest] - fn rejects_empty_directory() { - let dir = migrations_dir(&[]); - assert!(matches!(resolve(dir.path()), Err(Error::Empty { .. }))); - } -} diff --git a/litellm-rust/crates/migrate/Cargo.toml b/litellm-rust/crates/migrate/Cargo.toml deleted file mode 100644 index bb1ecaa3128..00000000000 --- a/litellm-rust/crates/migrate/Cargo.toml +++ /dev/null @@ -1,12 +0,0 @@ -[package] -name = "litellm-migrate" -version = "0.1.0" -edition.workspace = true -license.workspace = true -repository.workspace = true - -[dependencies] -litellm-migrate-macros.workspace = true - -[dev-dependencies] -rstest.workspace = true diff --git a/litellm-rust/crates/migrate/README.md b/litellm-rust/crates/migrate/README.md deleted file mode 100644 index 4817029451c..00000000000 --- a/litellm-rust/crates/migrate/README.md +++ /dev/null @@ -1,5 +0,0 @@ -# Migrations - -`litellm-migrate` exports the `Migration` struct and the `migrate!` macro that embeds a directory of `_.sql` files at compile time, sorted by numeric version - -The crate does not apply or track migrations; callers decide how and when the embedded SQL runs diff --git a/litellm-rust/crates/migrate/src/lib.rs b/litellm-rust/crates/migrate/src/lib.rs deleted file mode 100644 index f4e065e1b53..00000000000 --- a/litellm-rust/crates/migrate/src/lib.rs +++ /dev/null @@ -1,8 +0,0 @@ -pub use litellm_migrate_macros::migrate; - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub struct Migration { - pub version: u64, - pub description: &'static str, - pub sql: &'static str, -} diff --git a/litellm-rust/crates/migrate/tests/fixtures/migrations/10_tenth.sql b/litellm-rust/crates/migrate/tests/fixtures/migrations/10_tenth.sql deleted file mode 100644 index 31807719e9c..00000000000 --- a/litellm-rust/crates/migrate/tests/fixtures/migrations/10_tenth.sql +++ /dev/null @@ -1 +0,0 @@ -SELECT 10; diff --git a/litellm-rust/crates/migrate/tests/fixtures/migrations/1_first.sql b/litellm-rust/crates/migrate/tests/fixtures/migrations/1_first.sql deleted file mode 100644 index e0ac49d1ecf..00000000000 --- a/litellm-rust/crates/migrate/tests/fixtures/migrations/1_first.sql +++ /dev/null @@ -1 +0,0 @@ -SELECT 1; diff --git a/litellm-rust/crates/migrate/tests/fixtures/migrations/2_second.sql b/litellm-rust/crates/migrate/tests/fixtures/migrations/2_second.sql deleted file mode 100644 index e7f8100648d..00000000000 --- a/litellm-rust/crates/migrate/tests/fixtures/migrations/2_second.sql +++ /dev/null @@ -1 +0,0 @@ -SELECT 2; diff --git a/litellm-rust/crates/migrate/tests/migrate.rs b/litellm-rust/crates/migrate/tests/migrate.rs deleted file mode 100644 index 61c80351cf4..00000000000 --- a/litellm-rust/crates/migrate/tests/migrate.rs +++ /dev/null @@ -1,21 +0,0 @@ -use litellm_migrate::Migration; -use rstest::rstest; - -const MIGRATIONS: &[Migration] = litellm_migrate::migrate!("tests/fixtures/migrations"); - -#[rstest] -#[case::first(0, 1, "first", include_str!("fixtures/migrations/1_first.sql"))] -#[case::second(1, 2, "second", include_str!("fixtures/migrations/2_second.sql"))] -#[case::tenth(2, 10, "tenth", include_str!("fixtures/migrations/10_tenth.sql"))] -fn embeds_every_file_sorted_by_numeric_version( - #[case] index: usize, - #[case] version: u64, - #[case] description: &str, - #[case] sql: &str, -) { - assert_eq!(MIGRATIONS.len(), 3); - let migration = &MIGRATIONS[index]; - assert_eq!(migration.version, version); - assert_eq!(migration.description, description); - assert_eq!(migration.sql, sql); -} diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index 226cecd5b55..7458e64d374 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -74,6 +74,7 @@ serde_with.workspace = true criterion.workspace = true futures-util.workspace = true rstest.workspace = true +sqlx = { workspace = true, features = ["migrate"] } sha2.workspace = true tokio-tungstenite.workspace = true wiremock.workspace = true diff --git a/litellm-rust/crates/python-bridge/src/routes/traces.rs b/litellm-rust/crates/python-bridge/src/routes/traces.rs index f442d6c31de..65db1439d13 100644 --- a/litellm-rust/crates/python-bridge/src/routes/traces.rs +++ b/litellm-rust/crates/python-bridge/src/routes/traces.rs @@ -52,8 +52,7 @@ fn map_error_ref(error: &Error) -> PyErr { | Error::InvalidParameters | Error::InvalidScope => PyValueError::new_err(error.to_string()), Error::Task - | Error::SchemaFailed(_) - | Error::SchemaTransport + | Error::Migration(_) | Error::MissingSecret | Error::Busy | Error::ProvisionFailed(_) @@ -452,7 +451,14 @@ mod tests { )] #[case::insert_budget(Error::InsertTooLarge, "OverflowError")] #[case::scope(Error::InvalidScope, "ValueError")] - #[case::schema(Error::SchemaFailed(503), "RuntimeError")] + #[case::schema( + Error::Storage(litellm_storage_clickhouse::Error::SchemaFailed(503)), + "RuntimeError" + )] + #[case::migration( + Error::Migration(sqlx::migrate::MigrateError::VersionMismatch(1)), + "RuntimeError" + )] #[case::reader(Error::MissingSecret, "RuntimeError")] #[case::storage( Error::Storage(litellm_storage_clickhouse::Error::InvalidUrl), diff --git a/litellm-rust/crates/storage-clickhouse/AGENTS.md b/litellm-rust/crates/storage-clickhouse/AGENTS.md index 959ffdffb88..3f87e5d0759 100644 --- a/litellm-rust/crates/storage-clickhouse/AGENTS.md +++ b/litellm-rust/crates/storage-clickhouse/AGENTS.md @@ -3,3 +3,7 @@ `litellm-storage-clickhouse` exports `Storage`, a writer and bounded reader derived from one ClickHouse URL and database. It also exports bounded HTTP read and insert execution The crate has no trace tables, OTLP types, or named trace queries. `litellm-traces-clickhouse` supplies those rules and uses this storage for both trace rows and spend rows + +It also applies embedded SQLx migrations through the `_sqlx_migrations` ledger + +Migration files are append-only, and changed applied files are rejected by their checksums. Startup migrations must be replay-safe schema changes because the runner records success after execution without dirty states or locks. Backfills belong in coordinated jobs outside proxy startup. The replay policy lives in `ClickHouseMigrate::apply`, `dirty_version`, and `lock`; a Keeper-backed or deploy-time runner changes only those methods diff --git a/litellm-rust/crates/storage-clickhouse/Cargo.toml b/litellm-rust/crates/storage-clickhouse/Cargo.toml index acce941f2c9..4e72e823dc2 100644 --- a/litellm-rust/crates/storage-clickhouse/Cargo.toml +++ b/litellm-rust/crates/storage-clickhouse/Cargo.toml @@ -11,11 +11,14 @@ flate2.workspace = true litellm-http.workspace = true serde.workspace = true serde_json.workspace = true +serde_with.workspace = true +sqlx = { workspace = true, features = ["migrate"] } thiserror.workspace = true url.workspace = true [dev-dependencies] litellm-http = { workspace = true, features = ["test-support"] } rstest.workspace = true +testcontainers-modules = { version = "0.15.0", features = ["clickhouse"] } tokio.workspace = true wiremock.workspace = true diff --git a/litellm-rust/crates/storage-clickhouse/src/error.rs b/litellm-rust/crates/storage-clickhouse/src/error.rs index acec9b91675..fa0c5eabd73 100644 --- a/litellm-rust/crates/storage-clickhouse/src/error.rs +++ b/litellm-rust/crates/storage-clickhouse/src/error.rs @@ -1,4 +1,4 @@ -#[derive(Debug, thiserror::Error)] +#[derive(Clone, Debug, thiserror::Error)] pub enum Error { #[error("invalid ClickHouse insert row")] InvalidRow, diff --git a/litellm-rust/crates/storage-clickhouse/src/lib.rs b/litellm-rust/crates/storage-clickhouse/src/lib.rs index 7ab2aa9bc0a..e2ef62d1bc3 100644 --- a/litellm-rust/crates/storage-clickhouse/src/lib.rs +++ b/litellm-rust/crates/storage-clickhouse/src/lib.rs @@ -1,9 +1,11 @@ mod error; mod insert; +mod migrate; mod read; pub use error::Error; pub use insert::{insert_compressed_rows, insert_encoded_rows}; +pub use migrate::{ClickHouseMigrate, execute_statement, storage_error}; pub use read::{Parameter, Query, READ_LIMITS, ReadLimits, execute_read, fetch, fetch_json}; use url::Url; diff --git a/litellm-rust/crates/storage-clickhouse/src/migrate.rs b/litellm-rust/crates/storage-clickhouse/src/migrate.rs new file mode 100644 index 00000000000..c249020faaf --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/src/migrate.rs @@ -0,0 +1,308 @@ +use std::{ + future::Future, + pin::Pin, + time::{Duration, Instant}, +}; + +use litellm_http::Client; +use serde::Deserialize; +use serde_with::{DisplayFromStr, PickFirst, serde_as}; +use sqlx::{ + Error as SqlxError, + migrate::{AppliedMigration, Migrate, MigrateError, Migration}, +}; + +use crate::{Connection, Error, READ_LIMITS, valid_identifier}; + +pub async fn execute_statement( + client: &Client, + connection: &Connection, + sql: &str, + timeout: Duration, +) -> Result<(), Error> { + let body = execute_sql(client, connection, sql, timeout).await?; + if !body.trim().is_empty() { + return Err(Error::InvalidResponse); + } + Ok(()) +} + +async fn execute_sql( + client: &Client, + connection: &Connection, + sql: &str, + timeout: Duration, +) -> Result { + let mut url = connection.url().clone(); + let pairs: Vec<_> = url + .query_pairs() + .filter(|(key, _)| { + !matches!( + key.as_ref(), + "query" | "wait_end_of_query" | "send_progress_in_http_headers" | "async_insert" + ) + }) + .map(|(key, value)| (key.into_owned(), value.into_owned())) + .collect(); + url.query_pairs_mut() + .clear() + .extend_pairs(pairs) + .append_pair("wait_end_of_query", "1") + .append_pair("send_progress_in_http_headers", "0") + .append_pair("async_insert", "0"); + let mut response = client + .post(url) + .timeout(timeout) + .body(sql.to_owned()) + .send() + .await + .map_err(|_| Error::Transport)?; + if !response.status().is_success() { + return Err(Error::SchemaFailed(response.status().as_u16())); + } + let mut body = Vec::new(); + while let Some(chunk) = response.chunk().await.map_err(|_| Error::Transport)? { + if body.len() + chunk.len() > READ_LIMITS.response_bytes { + return Err(Error::ResponseTooLarge); + } + body.extend_from_slice(&chunk); + } + String::from_utf8(body).map_err(|_| Error::InvalidResponse) +} + +/// The startup runner records success after execution, never marks migrations dirty, and skips locks +/// A Keeper-backed or deploy-time runner changes only `apply`, `dirty_version`, and `lock` +pub struct ClickHouseMigrate<'a, R> { + client: &'a Client, + connection: &'a Connection, + database: &'a str, + render: R, + timeout: Duration, +} + +impl<'a, R> ClickHouseMigrate<'a, R> +where + R: Fn(&str) -> String + Send + Sync, +{ + pub fn new( + client: &'a Client, + connection: &'a Connection, + database: &'a str, + render: R, + timeout: Duration, + ) -> Result { + if !valid_identifier(database) { + return Err(Error::InvalidSchema); + } + Ok(Self { + client, + connection, + database, + render, + timeout, + }) + } +} + +#[serde_as] +#[derive(Deserialize)] +struct Applied { + #[serde_as(as = "PickFirst<(_, DisplayFromStr)>")] + version: i64, + checksum: String, +} + +fn migrate_error(error: Error) -> MigrateError { + MigrateError::Execute(SqlxError::AnyDriverError(Box::new(error))) +} + +fn migrate_execution_error(error: Error, version: i64) -> MigrateError { + MigrateError::ExecuteMigration(SqlxError::AnyDriverError(Box::new(error)), version) +} + +fn encode_hex(bytes: &[u8]) -> String { + bytes + .iter() + .map(|byte| format!("{byte:02x}")) + .collect::>() + .join("") +} + +fn decode_hex(value: &str) -> Result, Error> { + let (pairs, remainder) = value.as_bytes().as_chunks::<2>(); + if !remainder.is_empty() { + return Err(Error::InvalidResponse); + } + pairs + .iter() + .map(|pair| { + let high = decode_hex_digit(pair[0]).ok_or(Error::InvalidResponse)?; + let low = decode_hex_digit(pair[1]).ok_or(Error::InvalidResponse)?; + Ok((high << 4) | low) + }) + .collect() +} + +fn decode_hex_digit(value: u8) -> Option { + match value { + b'0'..=b'9' => Some(value - b'0'), + b'a'..=b'f' => Some(value - b'a' + 10), + b'A'..=b'F' => Some(value - b'A' + 10), + _ => None, + } +} + +fn escape_sql_string(value: &str) -> String { + value.replace('\\', "\\\\").replace('\'', "\\'") +} + +type MigrateFuture<'e, T> = Pin + Send + 'e>>; + +impl Migrate for ClickHouseMigrate<'_, R> +where + R: Fn(&str) -> String + Send + Sync, +{ + fn create_schema_if_not_exists<'e>( + &'e mut self, + schema_name: &'e str, + ) -> MigrateFuture<'e, Result<(), MigrateError>> { + Box::pin(async move { + if !valid_identifier(schema_name) { + return Err(migrate_error(Error::InvalidSchema)); + } + let statement = format!("CREATE DATABASE IF NOT EXISTS `{schema_name}`"); + execute_statement(self.client, self.connection, &statement, self.timeout) + .await + .map_err(migrate_error) + }) + } + + fn ensure_migrations_table<'e>( + &'e mut self, + table_name: &'e str, + ) -> MigrateFuture<'e, Result<(), MigrateError>> { + Box::pin(async move { + let database = format!("`{}`", self.database); + execute_statement( + self.client, + self.connection, + &format!("CREATE DATABASE IF NOT EXISTS {database}"), + self.timeout, + ) + .await + .map_err(migrate_error)?; + execute_statement( + self.client, + self.connection, + &format!( + "CREATE TABLE IF NOT EXISTS {database}.{table_name} \ + (version Int64, description String, installed_on DateTime64(3) DEFAULT now64(3), \ + success Bool, checksum String, execution_time Int64) ENGINE = MergeTree ORDER BY version" + ), + self.timeout, + ) + .await + .map_err(migrate_error) + }) + } + + fn dirty_version<'e>( + &'e mut self, + _table_name: &'e str, + ) -> MigrateFuture<'e, Result, MigrateError>> { + Box::pin(async { Ok(None) }) + } + + fn list_applied_migrations<'e>( + &'e mut self, + table_name: &'e str, + ) -> MigrateFuture<'e, Result, MigrateError>> { + Box::pin(async move { + let database = format!("`{}`", self.database); + let statement = format!( + "SELECT DISTINCT version, checksum FROM {database}.{table_name} \ + WHERE success ORDER BY version FORMAT JSONEachRow" + ); + let response = execute_sql(self.client, self.connection, &statement, self.timeout) + .await + .map_err(migrate_error)?; + response + .lines() + .filter(|line| !line.trim().is_empty()) + .map(|line| { + let row = serde_json::from_str::(line) + .map_err(|_| migrate_error(Error::InvalidResponse))?; + let checksum = decode_hex(&row.checksum).map_err(migrate_error)?; + Ok(AppliedMigration { + version: row.version, + checksum: checksum.into(), + }) + }) + .collect() + }) + } + + fn lock(&mut self) -> MigrateFuture<'_, Result<(), MigrateError>> { + Box::pin(async { Ok(()) }) + } + + fn unlock(&mut self) -> MigrateFuture<'_, Result<(), MigrateError>> { + Box::pin(async { Ok(()) }) + } + + fn apply<'e>( + &'e mut self, + table_name: &'e str, + migration: &'e Migration, + ) -> MigrateFuture<'e, Result> { + Box::pin(async move { + let started_at = Instant::now(); + let statement = (self.render)(migration.sql.as_str()); + execute_statement(self.client, self.connection, &statement, self.timeout) + .await + .map_err(|error| migrate_execution_error(error, migration.version))?; + let elapsed = started_at.elapsed(); + let execution_time = elapsed.as_nanos().min(i64::MAX as u128) as i64; + let description = escape_sql_string(&migration.description); + let checksum = encode_hex(&migration.checksum); + let database = format!("`{}`", self.database); + execute_statement( + self.client, + self.connection, + &format!( + "INSERT INTO {database}.{table_name} \ + (version, description, success, checksum, execution_time) \ + VALUES ({}, '{}', true, '{}', {execution_time})", + migration.version, description, checksum + ), + self.timeout, + ) + .await + .map_err(migrate_error)?; + Ok(elapsed) + }) + } + + fn revert<'e>( + &'e mut self, + _table_name: &'e str, + _migration: &'e Migration, + ) -> MigrateFuture<'e, Result> { + Box::pin(async { + Err(MigrateError::Execute(SqlxError::AnyDriverError(Box::new( + std::io::Error::other("ClickHouse migrations are forward-only"), + )))) + }) + } +} + +pub fn storage_error(error: &MigrateError) -> Option<&Error> { + let error = match error { + MigrateError::Execute(error) | MigrateError::ExecuteMigration(error, _) => error, + _ => return None, + }; + match error { + SqlxError::AnyDriverError(error) => error.downcast_ref(), + _ => None, + } +} diff --git a/litellm-rust/crates/storage-clickhouse/tests/migrations.rs b/litellm-rust/crates/storage-clickhouse/tests/migrations.rs new file mode 100644 index 00000000000..29a85a8c4ec --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/tests/migrations.rs @@ -0,0 +1,476 @@ +use std::time::Duration; + +use litellm_http::Client; +use litellm_storage_clickhouse::{ + ClickHouseMigrate, Connection, Error, READ_LIMITS, execute_statement, storage_error, +}; +use rstest::{fixture, rstest}; +use sqlx::{ + SqlStr, + migrate::{Migrate, MigrateError, Migration, MigrationType, Migrator}, +}; +use testcontainers_modules::{ + clickhouse::ClickHouse, + testcontainers::{ContainerAsync, ImageExt, runners::AsyncRunner}, +}; +use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{body_string, method, query_param}, +}; + +const CLICKHOUSE_TAG: &str = + "26.9.6.6@sha256:eb4870e7ca7ed70c259eebfcfbee6cf797017f6b5436c2926bbbfe3d4d28486e"; +const DATABASE: &str = "storage_migrate_test"; +const REQUEST_TIMEOUT: Duration = Duration::from_secs(10); +const SELECT_APPLIED: &str = "SELECT DISTINCT version, checksum FROM `trace_test`._sqlx_migrations \ + WHERE success ORDER BY version FORMAT JSONEachRow"; + +type TestResult = Result>; + +struct ClickHouseDatabase { + _container: ContainerAsync, + url: String, + client: Client, +} + +#[fixture] +async fn database() -> TestResult { + let container = ClickHouse::default() + .with_tag(CLICKHOUSE_TAG) + .with_env_var("CLICKHOUSE_SKIP_USER_SETUP", "1") + .start() + .await?; + let url = format!( + "http://{}:{}", + container.get_host().await?, + container.get_host_port_ipv4(8123).await? + ); + Ok(ClickHouseDatabase { + _container: container, + url, + client: Client::no_redirect_for_test(), + }) +} + +#[fixture] +async fn mock_server() -> MockServer { + let server = MockServer::start().await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(200)) + .with_priority(10) + .mount(&server) + .await; + server +} + +fn migration(version: i64, sql: &'static str) -> Migration { + Migration::new( + version, + format!("migration_{version}").into(), + MigrationType::Simple, + SqlStr::from_static(sql), + false, + ) +} + +fn migrator(migrations: Vec) -> Migrator { + Migrator { + ignore_missing: true, + locking: false, + ..Migrator::with_migrations(migrations) + } +} + +fn render_database(sql: &str) -> String { + sql.replace("{database}", &format!("`{DATABASE}`")) +} + +async fn run_migrations( + database: &ClickHouseDatabase, + migrator: &Migrator, + schema: &str, + render: R, +) -> Result<(), MigrateError> +where + R: Fn(&str) -> String + Send + Sync, +{ + let connection = Connection::writer(&database.url).expect("valid ClickHouse URL"); + let mut adapter = ClickHouseMigrate::new( + &database.client, + &connection, + schema, + render, + REQUEST_TIMEOUT, + ) + .expect("valid schema"); + migrator.run_direct(None, &mut adapter, false).await +} + +async fn execute_write(database: &ClickHouseDatabase, sql: &str) -> TestResult { + database + .client + .post(&database.url) + .body(sql.to_owned()) + .send() + .await? + .error_for_status()?; + Ok(()) +} + +async fn read_json(database: &ClickHouseDatabase, sql: &str) -> TestResult { + let response = database + .client + .post(&database.url) + .body(sql.to_owned()) + .send() + .await? + .error_for_status()?; + Ok(serde_json::from_str(&response.text().await?)?) +} + +async fn ledger_versions(database: &ClickHouseDatabase) -> TestResult> { + let response = read_json( + database, + &format!( + "SELECT version FROM `{DATABASE}`._sqlx_migrations \ + GROUP BY version ORDER BY version FORMAT JSON" + ), + ) + .await?; + Ok(response["data"] + .as_array() + .expect("ClickHouse returns versions") + .iter() + .map(|row| row["version"].as_i64().expect("version is Int64")) + .collect()) +} + +fn encode_hex(bytes: &[u8]) -> String { + bytes + .iter() + .map(|byte| format!("{byte:02x}")) + .collect::>() + .join("") +} + +#[rstest] +#[tokio::test] +async fn only_pending_migrations_execute_on_the_second_run( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let migrator = migrator(vec![ + migration( + 1, + "CREATE TABLE IF NOT EXISTS {database}.migration_one (id UInt8) ENGINE = MergeTree ORDER BY id", + ), + migration( + 2, + "CREATE TABLE IF NOT EXISTS {database}.migration_two (id UInt8) ENGINE = MergeTree ORDER BY id", + ), + ]); + run_migrations(&database, &migrator, DATABASE, render_database).await?; + run_migrations(&database, &migrator, DATABASE, render_database).await?; + execute_statement( + &database.client, + &Connection::writer(&database.url)?, + "SYSTEM FLUSH LOGS", + REQUEST_TIMEOUT, + ) + .await?; + + let queries = read_json( + &database, + "SELECT count() AS executions FROM system.query_log \ + WHERE type = 'QueryFinish' AND query LIKE \ + 'CREATE TABLE IF NOT EXISTS `storage_migrate_test`.migration_%' FORMAT JSON", + ) + .await?; + assert_eq!(queries["data"][0]["executions"].as_u64(), Some(2)); + assert_eq!(ledger_versions(&database).await?, vec![1, 2]); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn edited_migration_checksum_returns_version_mismatch( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let original = migrator(vec![migration( + 1, + "CREATE TABLE IF NOT EXISTS {database}.original (id UInt8) ENGINE = MergeTree ORDER BY id", + )]); + let changed = migrator(vec![migration( + 1, + "CREATE TABLE IF NOT EXISTS {database}.changed (id UInt8) ENGINE = MergeTree ORDER BY id", + )]); + run_migrations(&database, &original, DATABASE, render_database).await?; + + assert!(matches!( + run_migrations(&database, &changed, DATABASE, render_database).await, + Err(MigrateError::VersionMismatch(1)) + )); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn failed_migration_is_not_recorded_and_retains_storage_error( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let migrator = migrator(vec![migration(1, "THIS IS NOT VALID CLICKHOUSE SQL")]); + let error = run_migrations(&database, &migrator, DATABASE, render_database) + .await + .expect_err("invalid SQL must fail"); + assert!(matches!(&error, MigrateError::ExecuteMigration(_, 1))); + assert!(matches!( + storage_error(&error), + Some(Error::SchemaFailed(_)) + )); + + let rows = read_json( + &database, + &format!( + "SELECT count() AS rows FROM `{DATABASE}`._sqlx_migrations \ + WHERE version = 1 FORMAT JSON" + ), + ) + .await?; + assert_eq!(rows["data"][0]["rows"].as_u64(), Some(0)); + Ok(()) +} + +#[rstest] +fn invalid_database_identifier_is_rejected() { + let client = Client::no_redirect_for_test(); + let connection = Connection::writer("http://127.0.0.1:1").expect("valid URL"); + assert!(matches!( + ClickHouseMigrate::new( + &client, + &connection, + "storage_test; DROP DATABASE default", + str::to_owned, + REQUEST_TIMEOUT, + ), + Err(Error::InvalidSchema) + )); +} + +#[rstest] +#[tokio::test] +async fn unknown_source_version_is_tolerated( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + run_migrations(&database, &migrator(vec![]), DATABASE, render_database).await?; + execute_write( + &database, + &format!( + "INSERT INTO `{DATABASE}`._sqlx_migrations \ + (version, description, success, checksum, execution_time) \ + VALUES (99, 'unknown', true, '{}', 0)", + "00".repeat(48) + ), + ) + .await?; + let migrator = migrator(vec![migration( + 1, + "CREATE TABLE IF NOT EXISTS {database}.known (id UInt8) ENGINE = MergeTree ORDER BY id", + )]); + + run_migrations(&database, &migrator, DATABASE, render_database).await?; + + assert_eq!(ledger_versions(&database).await?, vec![1, 99]); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn duplicate_ledger_rows_are_tolerated( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let applied = migration( + 1, + "CREATE TABLE IF NOT EXISTS {database}.duplicate_test (id UInt8) ENGINE = MergeTree ORDER BY id", + ); + let checksum = encode_hex(&applied.checksum); + let migrator = migrator(vec![applied]); + run_migrations(&database, &migrator, DATABASE, render_database).await?; + execute_write( + &database, + &format!( + "INSERT INTO `{DATABASE}`._sqlx_migrations \ + (version, description, success, checksum, execution_time) \ + VALUES (1, 'migration_1', true, '{checksum}', 0)" + ), + ) + .await?; + + run_migrations(&database, &migrator, DATABASE, render_database).await?; + + let rows = read_json( + &database, + &format!( + "SELECT count() AS rows, uniqExact(version) AS versions \ + FROM `{DATABASE}`._sqlx_migrations FORMAT JSON" + ), + ) + .await?; + assert_eq!(rows["data"][0]["rows"].as_u64(), Some(2)); + assert_eq!(rows["data"][0]["versions"].as_u64(), Some(1)); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn revert_reports_forward_only_error() { + let client = Client::no_redirect_for_test(); + let connection = Connection::writer("http://127.0.0.1:1").expect("valid URL"); + let migration = migration( + 1, + "CREATE TABLE IF NOT EXISTS {database}.revert_test (id UInt8) ENGINE = MergeTree ORDER BY id", + ); + let mut adapter = ClickHouseMigrate::new( + &client, + &connection, + DATABASE, + render_database, + REQUEST_TIMEOUT, + ) + .expect("valid schema"); + + let error = adapter + .revert("_sqlx_migrations", &migration) + .await + .expect_err("ClickHouse migrations cannot be reverted"); + assert!( + error + .to_string() + .contains("ClickHouse migrations are forward-only") + ); +} + +#[rstest] +#[case::numeric("1")] +#[case::quoted("\"1\"")] +#[tokio::test] +async fn applied_int64_versions_accept_numeric_and_quoted_json( + #[future(awt)] mock_server: MockServer, + #[case] version: &str, +) { + Mock::given(method("POST")) + .and(body_string(SELECT_APPLIED)) + .respond_with( + ResponseTemplate::new(200) + .set_body_string(format!("{{\"version\":{version},\"checksum\":\"00\"}}\n")), + ) + .mount(&mock_server) + .await; + let client = Client::no_redirect_for_test(); + let connection = Connection::writer(&mock_server.uri()).expect("valid URL"); + let migrator = migrator(vec![]); + let mut adapter = ClickHouseMigrate::new( + &client, + &connection, + "trace_test", + str::to_owned, + REQUEST_TIMEOUT, + ) + .expect("valid schema"); + + migrator + .run_direct(None, &mut adapter, false) + .await + .expect("applied version parses"); +} + +#[rstest] +#[tokio::test] +async fn schema_requests_override_unsafe_connection_settings( + #[future(awt)] mock_server: MockServer, +) { + Mock::given(method("POST")) + .and(query_param("wait_end_of_query", "1")) + .and(query_param("send_progress_in_http_headers", "0")) + .and(query_param("async_insert", "0")) + .and(query_param("custom_setting", "preserved")) + .respond_with(ResponseTemplate::new(200)) + .expect(1) + .mount(&mock_server) + .await; + let url = format!( + "{}/?wait_end_of_query=0&send_progress_in_http_headers=1&async_insert=1&custom_setting=preserved", + mock_server.uri() + ); + execute_statement( + &Client::no_redirect_for_test(), + &Connection::writer(&url).expect("valid URL"), + "CREATE DATABASE IF NOT EXISTS trace_test", + REQUEST_TIMEOUT, + ) + .await + .expect("schema execution succeeds"); + let requests = mock_server + .received_requests() + .await + .expect("requests recorded"); + for name in [ + "wait_end_of_query", + "send_progress_in_http_headers", + "async_insert", + ] { + assert_eq!( + requests[0] + .url + .query_pairs() + .filter(|(key, _)| key == name) + .count(), + 1 + ); + } +} + +#[rstest] +#[tokio::test] +async fn oversized_ledger_response_is_rejected_before_migrations( + #[future(awt)] mock_server: MockServer, +) { + Mock::given(method("POST")) + .and(body_string(SELECT_APPLIED)) + .respond_with( + ResponseTemplate::new(200).set_body_string(" ".repeat(READ_LIMITS.response_bytes + 1)), + ) + .mount(&mock_server) + .await; + let client = Client::no_redirect_for_test(); + let connection = Connection::writer(&mock_server.uri()).expect("valid URL"); + let migrator = migrator(vec![]); + let mut adapter = ClickHouseMigrate::new( + &client, + &connection, + "trace_test", + str::to_owned, + REQUEST_TIMEOUT, + ) + .expect("valid schema"); + let error = migrator + .run_direct(None, &mut adapter, false) + .await + .expect_err("oversized result is rejected"); + + assert!(matches!( + storage_error(&error), + Some(Error::ResponseTooLarge) + )); + assert_eq!( + mock_server + .received_requests() + .await + .expect("requests recorded") + .len(), + 3 + ); +} diff --git a/litellm-rust/crates/traces-clickhouse/AGENTS.md b/litellm-rust/crates/traces-clickhouse/AGENTS.md index ae0c1c8eeb9..462100b10a0 100644 --- a/litellm-rust/crates/traces-clickhouse/AGENTS.md +++ b/litellm-rust/crates/traces-clickhouse/AGENTS.md @@ -1,6 +1,7 @@ - Own trace schema, row encoding, SQL query adapters and reader provisioning; consume domain types from `litellm-traces` - Keep generic ClickHouse connections and HTTP execution in `litellm-storage-clickhouse`; keep PyO3 conversion in `python-bridge` -- Keep schema definitions only in `migrations/NNNN_description.sql`, embedded by `litellm_migrate::migrate!` +- Keep schema definitions only in `migrations/NNNN_description.sql`, embedded by `sqlx::migrate!` +- Treat retention TTLs as current configuration: change them in `RETENTION` in `src/schema.rs`, which every startup reapplies, never in a new migration - Require typed query parameters and SELECT-only readers with server-side limits and tenant isolation - Bound insert time and encoded bytes; preserve shared values and explicit retry deduplication - Test storage behavior through the public API against ClickHouse diff --git a/litellm-rust/crates/traces-clickhouse/Cargo.toml b/litellm-rust/crates/traces-clickhouse/Cargo.toml index b90ad8a7bf8..13aa6d06af6 100644 --- a/litellm-rust/crates/traces-clickhouse/Cargo.toml +++ b/litellm-rust/crates/traces-clickhouse/Cargo.toml @@ -16,7 +16,6 @@ flate2.workspace = true futures-util.workspace = true hmac = "0.12.1" litellm-http.workspace = true -litellm-migrate.workspace = true litellm-storage-clickhouse.workspace = true litellm-traces.workspace = true litellm-traces-cache.workspace = true @@ -24,6 +23,7 @@ moka.workspace = true serde.workspace = true serde_json.workspace = true sha2.workspace = true +sqlx = { workspace = true, features = ["migrate", "macros"] } strum.workspace = true thiserror.workspace = true time = { workspace = true, features = ["formatting"] } diff --git a/litellm-rust/crates/traces-clickhouse/src/error.rs b/litellm-rust/crates/traces-clickhouse/src/error.rs index 1d15c556316..e0ef127a0e9 100644 --- a/litellm-rust/crates/traces-clickhouse/src/error.rs +++ b/litellm-rust/crates/traces-clickhouse/src/error.rs @@ -16,10 +16,6 @@ pub enum Error { InvalidResponse, #[error("ClickHouse insert exceeds the encoded size limit")] InsertTooLarge, - #[error("ClickHouse schema setup failed with HTTP status {0}")] - SchemaFailed(u16), - #[error("ClickHouse schema setup transport failed")] - SchemaTransport, #[error("trace SQL queries require a configured proxy master key")] MissingSecret, #[error("invalid trace query scope")] @@ -39,5 +35,7 @@ pub enum Error { #[error(transparent)] Storage(#[from] litellm_storage_clickhouse::Error), #[error(transparent)] + Migration(#[from] sqlx::migrate::MigrateError), + #[error(transparent)] Cached(#[from] std::sync::Arc), } diff --git a/litellm-rust/crates/traces-clickhouse/src/lib.rs b/litellm-rust/crates/traces-clickhouse/src/lib.rs index 43ff0b8bf33..aed87691958 100644 --- a/litellm-rust/crates/traces-clickhouse/src/lib.rs +++ b/litellm-rust/crates/traces-clickhouse/src/lib.rs @@ -33,7 +33,8 @@ pub use query::{QueryHelp, execute_read, query_help, query_sql}; pub use query_access::QueryReaders; pub use reads::ClickHouseTraces; pub use schema::{ - NORMALIZED_FIELD_DEFINITIONS, NormalizedFieldDefinition, ensure_schema, schema_statements, + NORMALIZED_FIELD_DEFINITIONS, NormalizedFieldDefinition, apply_migrations, ensure_schema, + reconcile_retention, schema_statements, }; pub use span_row::span_rows; pub use sql::execute_named_read; diff --git a/litellm-rust/crates/traces-clickhouse/src/schema.rs b/litellm-rust/crates/traces-clickhouse/src/schema.rs index 6a39bd24041..562dbb976c2 100644 --- a/litellm-rust/crates/traces-clickhouse/src/schema.rs +++ b/litellm-rust/crates/traces-clickhouse/src/schema.rs @@ -1,15 +1,26 @@ use litellm_http::Client; -use litellm_migrate::Migration; +use litellm_storage_clickhouse::{ClickHouseMigrate, execute_statement, storage_error}; use serde::Serialize; +use sqlx::migrate::Migrator; use std::time::Duration; use super::{Connection, Error}; const SCHEMA_REQUEST_TIMEOUT: Duration = Duration::from_secs(30); -const MIGRATIONS: &[Migration] = litellm_migrate::migrate!("migrations"); +static MIGRATOR: Migrator = Migrator { + ignore_missing: true, + locking: false, + ..sqlx::migrate!("./migrations") +}; -pub fn schema_statements(database: &str, retention_days: u32) -> Result, Error> { +const RETENTION: [(&str, &str); 3] = [ + ("otel_traces", "toDateTime(Timestamp)"), + ("agent_traces_by_key", "toDateTime(StartTs)"), + ("spend_logs", "toDateTime(start_time)"), +]; + +fn validate_schema(database: &str, retention_days: u32) -> Result<(), Error> { if database.is_empty() || !database .bytes() @@ -18,19 +29,109 @@ pub fn schema_statements(database: &str, retention_days: u32) -> Result String { + sql.replace("{database}", database) + .replace("{retention_days}", &retention_days.to_string()) +} + +pub fn schema_statements(database: &str, retention_days: u32) -> Result, Error> { + validate_schema(database, retention_days)?; let database = format!("`{database}`"); Ok( std::iter::once(format!("CREATE DATABASE IF NOT EXISTS {database}")) - .chain(MIGRATIONS.iter().map(|migration| { - migration - .sql - .replace("{database}", &database) - .replace("{retention_days}", &retention_days.to_string()) - })) + .chain( + MIGRATOR + .migrations + .iter() + .map(|migration| render(migration.sql.as_str(), &database, retention_days)), + ) .collect(), ) } +pub async fn apply_migrations( + client: &Client, + connection: &Connection, + database: &str, + retention_days: u32, +) -> Result<(), Error> { + apply_migrations_with_timeout( + client, + connection, + database, + retention_days, + SCHEMA_REQUEST_TIMEOUT, + ) + .await +} + +async fn apply_migrations_with_timeout( + client: &Client, + connection: &Connection, + database: &str, + retention_days: u32, + request_timeout: Duration, +) -> Result<(), Error> { + validate_schema(database, retention_days)?; + let quoted_database = format!("`{database}`"); + let mut adapter = ClickHouseMigrate::new( + client, + connection, + database, + |sql| render(sql, "ed_database, retention_days), + request_timeout, + )?; + MIGRATOR + .run_direct(None, &mut adapter, false) + .await + .map_err(|error| match storage_error(&error) { + Some(storage_error) => Error::Storage(storage_error.clone()), + None => Error::Migration(error), + }) +} + +pub async fn reconcile_retention( + client: &Client, + connection: &Connection, + database: &str, + retention_days: u32, +) -> Result<(), Error> { + reconcile_retention_with_timeout( + client, + connection, + database, + retention_days, + SCHEMA_REQUEST_TIMEOUT, + ) + .await +} + +async fn reconcile_retention_with_timeout( + client: &Client, + connection: &Connection, + database: &str, + retention_days: u32, + request_timeout: Duration, +) -> Result<(), Error> { + validate_schema(database, retention_days)?; + let database = format!("`{database}`"); + for (table, expression) in RETENTION { + execute_statement( + client, + connection, + &format!( + "ALTER TABLE {database}.{table} MODIFY TTL {expression} + INTERVAL {retention_days} DAY" + ), + request_timeout, + ) + .await?; + } + Ok(()) +} + pub async fn ensure_schema( client: &Client, connection: &Connection, @@ -54,19 +155,22 @@ async fn ensure_schema_with_timeout( retention_days: u32, request_timeout: Duration, ) -> Result<(), Error> { - for statement in schema_statements(database, retention_days)? { - let response = client - .post(connection.url().clone()) - .timeout(request_timeout) - .body(statement) - .send() - .await - .map_err(|_| Error::SchemaTransport)?; - if !response.status().is_success() { - return Err(Error::SchemaFailed(response.status().as_u16())); - } - } - Ok(()) + apply_migrations_with_timeout( + client, + connection, + database, + retention_days, + request_timeout, + ) + .await?; + reconcile_retention_with_timeout( + client, + connection, + database, + retention_days, + request_timeout, + ) + .await } #[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize)] diff --git a/litellm-rust/crates/traces-clickhouse/tests/migrations.rs b/litellm-rust/crates/traces-clickhouse/tests/migrations.rs index 09fd3e78da8..8976ef208d8 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/migrations.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/migrations.rs @@ -1,11 +1,14 @@ use std::{collections::BTreeMap, time::Duration}; use litellm_http::Client; +use litellm_storage_clickhouse::Error as StorageError; use litellm_traces_clickhouse::{ Connection, Error, InsertTable, NORMALIZED_FIELD_DEFINITIONS, Parameter, ReadQuery, - encode_rows, ensure_schema, execute_named_read, execute_read, schema_statements, + apply_migrations, encode_rows, ensure_schema, execute_named_read, execute_read, + reconcile_retention, schema_statements, }; use rstest::rstest; +use sqlx::migrate::MigrateError; mod support; use support::{ClickHouseDatabase, TestResult, database}; @@ -71,6 +74,44 @@ async fn mutation_rows(database: &ClickHouseDatabase) -> TestResult { .expect("ClickHouse returns mutation counts as unsigned integers")) } +fn migration_versions() -> Vec { + let mut versions = std::fs::read_dir(concat!(env!("CARGO_MANIFEST_DIR"), "/migrations")) + .expect("migration directory exists") + .map(|entry| { + entry + .expect("migration directory entry is readable") + .file_name() + .into_string() + .expect("migration file name is UTF-8") + }) + .filter_map(|name| { + name.strip_suffix(".sql") + .and_then(|stem| stem.split('_').next()) + .and_then(|version| version.parse::().ok()) + }) + .collect::>(); + versions.sort_unstable(); + versions +} + +async fn migration_ledger_versions(database: &ClickHouseDatabase) -> TestResult> { + let response = read_json( + database, + "SELECT version FROM trace_test._sqlx_migrations GROUP BY version ORDER BY version", + ) + .await?; + Ok(response["data"] + .as_array() + .expect("ClickHouse returns version rows") + .iter() + .map(|row| { + row["version"] + .as_u64() + .expect("ClickHouse returns versions as unsigned integers") + }) + .collect()) +} + #[rstest] #[tokio::test] async fn schema_supports_span_rollups_and_spend_joins( @@ -80,6 +121,27 @@ async fn schema_supports_span_rollups_and_spend_joins( let writer = Connection::writer(&database.url)?; ensure_schema(&database.client, &writer, "trace_test", 7).await?; ensure_schema(&database.client, &writer, "trace_test", 7).await?; + let expected_versions = migration_versions(); + let ledger = read_json( + &database, + "SELECT count() AS rows, uniqExact(version) AS versions, \ + countIf(NOT match(checksum, '^[0-9a-f]{96}$')) AS invalid_checksums \ + FROM trace_test._sqlx_migrations", + ) + .await?; + assert_eq!( + ledger["data"][0]["rows"].as_u64(), + Some(expected_versions.len() as u64) + ); + assert_eq!( + ledger["data"][0]["versions"].as_u64(), + Some(expected_versions.len() as u64) + ); + assert_eq!(ledger["data"][0]["invalid_checksums"].as_u64(), Some(0)); + assert_eq!( + migration_ledger_versions(&database).await?, + expected_versions + ); let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; let span = serde_json::from_value(serde_json::json!({ "Timestamp": timestamp, "TraceId": "trace-1", "SpanId": "span-1", "ParentSpanId": "", @@ -202,6 +264,112 @@ async fn schema_supports_span_rollups_and_spend_joins( Ok(()) } +#[rstest] +#[case::quoted_versions("output_format_json_quote_64bit_integers=1")] +#[case::asynchronous_inserts( + "async_insert=1&wait_for_async_insert=0&async_insert_busy_timeout_ms=20000" +)] +#[tokio::test] +async fn schema_setup_records_migrations_synchronously_with_configured_settings( + #[future(awt)] database: TestResult, + #[case] settings: &str, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&format!("{}?{settings}", database.url))?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; + assert_eq!( + migration_ledger_versions(&database).await?, + migration_versions() + ); + ensure_schema(&database.client, &writer, "trace_test", 7).await?; + assert_eq!( + migration_ledger_versions(&database).await?, + migration_versions() + ); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn changed_migration_is_rejected( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&database.url)?; + ensure_schema(&database.client, &writer, "trace_test", 7).await?; + execute_write( + &database, + "ALTER TABLE trace_test._sqlx_migrations UPDATE checksum = '00' \ + WHERE version = 1 SETTINGS mutations_sync = 1", + ) + .await?; + + assert!(matches!( + ensure_schema(&database.client, &writer, "trace_test", 7).await, + Err(Error::Migration(MigrateError::VersionMismatch(1))) + )); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn concurrent_schema_setup_succeeds( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&database.url)?; + let (first, second, third, fourth) = tokio::join!( + ensure_schema(&database.client, &writer, "trace_test", 7), + ensure_schema(&database.client, &writer, "trace_test", 7), + ensure_schema(&database.client, &writer, "trace_test", 7), + ensure_schema(&database.client, &writer, "trace_test", 7), + ); + for result in [first, second, third, fourth] { + result?; + } + let tables = read_json( + &database, + "SELECT count() AS tables FROM system.tables \ + WHERE database = 'trace_test' AND name IN \ + ('otel_traces', 'agent_traces_by_key', 'spend_logs')", + ) + .await?; + assert_eq!(tables["data"][0]["tables"].as_u64(), Some(3)); + assert_eq!( + migration_ledger_versions(&database).await?, + migration_versions() + ); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn existing_schema_without_ledger_is_adopted( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&database.url)?; + for statement in schema_statements("trace_test", 7)? { + execute_write(&database, &statement).await?; + } + let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; + let span = serde_json::from_value(serde_json::json!({ + "Timestamp": timestamp, "TraceId": "trace-adopted", "SpanId": "span-adopted", + "ParentSpanId": "", "ServiceName": "proxy", "SpanName": "request", + "Input": "existing row", "ResourceAttributes": {}, "SpanAttributes": {} + }))?; + insert_rows(&database, "otel_traces", vec![span]).await?; + + ensure_schema(&database.client, &writer, "trace_test", 7).await?; + + assert_eq!(table_rows(&database, "otel_traces").await?, 1); + assert_eq!( + migration_ledger_versions(&database).await?, + migration_versions() + ); + Ok(()) +} + #[rstest] #[tokio::test] async fn normalized_fields_match_clickhouse_catalog( @@ -785,6 +953,60 @@ async fn retention_changes_materialize_existing_rows_and_remain_idempotent( Ok(()) } +#[rstest] +#[tokio::test] +async fn retention_reconciliation_updates_each_table_ttl( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&database.url)?; + apply_migrations(&database.client, &writer, "trace_test", 7).await?; + reconcile_retention(&database.client, &writer, "trace_test", 7).await?; + let ttl_queries = read_json( + &database, + "SELECT name, create_table_query FROM system.tables \ + WHERE database = 'trace_test' AND name IN \ + ('otel_traces', 'agent_traces_by_key', 'spend_logs') ORDER BY name", + ) + .await?; + let ttl_queries = ttl_queries["data"].as_array().expect("retention tables"); + assert_eq!( + ttl_queries + .iter() + .map(|row| row["name"].as_str().expect("table name")) + .collect::>(), + ["agent_traces_by_key", "otel_traces", "spend_logs"] + ); + for row in ttl_queries { + let query = row["create_table_query"] + .as_str() + .expect("table creation query"); + assert!( + query.contains("toIntervalDay(7)") || query.contains("INTERVAL 7 DAY"), + "{query}" + ); + } + + reconcile_retention(&database.client, &writer, "trace_test", 3).await?; + let ttl_queries = read_json( + &database, + "SELECT name, create_table_query FROM system.tables \ + WHERE database = 'trace_test' AND name IN \ + ('otel_traces', 'agent_traces_by_key', 'spend_logs') ORDER BY name", + ) + .await?; + for row in ttl_queries["data"].as_array().expect("retention tables") { + let query = row["create_table_query"] + .as_str() + .expect("table creation query"); + assert!( + query.contains("toIntervalDay(3)") || query.contains("INTERVAL 3 DAY"), + "{query}" + ); + } + Ok(()) +} + #[rstest] #[tokio::test] async fn schema_statement_timeout_maps_to_transport_error() -> TestResult { @@ -804,7 +1026,7 @@ async fn schema_statement_timeout_maps_to_transport_error() -> TestResult { .await; server.abort(); assert!( - matches!(result, Ok(Err(Error::SchemaTransport))), + matches!(result, Ok(Err(Error::Storage(StorageError::Transport)))), "{result:?}" ); Ok(()) diff --git a/tests/test_litellm_rust/test_traces.py b/tests/test_litellm_rust/test_traces.py index b8af0f0c374..43950a2f011 100644 --- a/tests/test_litellm_rust/test_traces.py +++ b/tests/test_litellm_rust/test_traces.py @@ -167,13 +167,18 @@ async def test_from_env_reads_with_clickhouse_url( @pytest.mark.asyncio async def test_schema_setup_uses_configured_retention(recording_server: RecordingServer) -> None: recording_server.expected_requests = None + recording_server.default_response = ResponseSpec(body=b"") storage: Final = _native_storage("trace_test", recording_server.base_url, 7) await storage.ensure_schema() ttl_statements: Final = tuple( request.raw_body for request in recording_server.requests if b"MODIFY TTL" in request.raw_body ) - assert len(ttl_statements) == 3 assert all(b"INTERVAL 7 DAY" in statement for statement in ttl_statements) + assert tuple(request.raw_body.strip() for request in recording_server.requests[-3:]) == ( + b"ALTER TABLE `trace_test`.otel_traces MODIFY TTL toDateTime(Timestamp) + INTERVAL 7 DAY", + b"ALTER TABLE `trace_test`.agent_traces_by_key MODIFY TTL toDateTime(StartTs) + INTERVAL 7 DAY", + b"ALTER TABLE `trace_test`.spend_logs MODIFY TTL toDateTime(start_time) + INTERVAL 7 DAY", + ) @pytest.mark.asyncio @@ -181,7 +186,7 @@ async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement recording_server: RecordingServer, ) -> None: recording_server.expected_requests = 2 - recording_server.enqueue(ResponseSpec(body="")) + recording_server.enqueue(ResponseSpec(body=b"")) recording_server.enqueue(ResponseSpec(status=403, body="denied")) writer_url: Final = recording_server.base_url.replace("http://", "http://writer:p%40ss%2Fword%25@") storage: Final = _native_storage("trace_test", writer_url + "?database=wrong&readonly=1", 7) From 21881c571181fc0e409dd717b8a277e5b43152a7 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 5 Oct 2026 16:11:52 +0000 Subject: [PATCH 08/19] refactor(ui): extract shared timeline and time-range controls (#44584) Move the timeline renderer to shared/timeline/Timeline taking buckets, a selected window, and callbacks as props, and TimeRangeControls to shared/timeline. Lens keeps bucketRuns as the adapter that converts loaded traces into buckets, preserving behavior without the histogram endpoint. Co-authored-by: Yujong Lee Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../lens/traces/list/AgentTracesSection.tsx | 5 +- .../lens/traces/list/TracesTimeline.test.ts | 42 +-- .../lens/traces/list/TracesTimeline.tsx | 318 +---------------- .../lens/traces/list/runSearch/RunSearch.tsx | 2 +- .../traces/list/runSearch/RunsToolbar.tsx | 2 +- .../lens/traces/list/runSearch/runSql.ts | 2 +- .../src/components/lens/traces/routing.ts | 4 +- .../timeline}/TimeRangeControls.tsx | 2 +- .../shared/timeline/Timeline.test.ts | 55 +++ .../components/shared/timeline/Timeline.tsx | 332 ++++++++++++++++++ 10 files changed, 402 insertions(+), 362 deletions(-) rename ui/litellm-dashboard/src/components/{lens/traces/list => shared/timeline}/TimeRangeControls.tsx (98%) create mode 100644 ui/litellm-dashboard/src/components/shared/timeline/Timeline.test.ts create mode 100644 ui/litellm-dashboard/src/components/shared/timeline/Timeline.tsx diff --git a/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesSection.tsx b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesSection.tsx index ffdb3a959a5..4f18584d9ce 100644 --- a/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesSection.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/list/AgentTracesSection.tsx @@ -19,8 +19,9 @@ import { } from "../routing"; import type { TraceSummary } from "../types"; import { RunView } from "../detail/run/RunView"; -import { TimeRangeControls } from "./TimeRangeControls"; -import { TracesTimeline, type TimeWindow } from "./TracesTimeline"; +import { TimeRangeControls } from "@/components/shared/timeline/TimeRangeControls"; +import { TracesTimeline } from "./TracesTimeline"; +import type { TimeWindow } from "@/components/shared/timeline/Timeline"; import { TracingSetupCard } from "../../onboarding/tracing/TracingSetupCard"; import { useTracesLive } from "../api"; import { type AgentTracesResult, traceWindowStartMs, useAgentTraces, useTraceAvailability } from "./useAgentTraces"; diff --git a/ui/litellm-dashboard/src/components/lens/traces/list/TracesTimeline.test.ts b/ui/litellm-dashboard/src/components/lens/traces/list/TracesTimeline.test.ts index 4da827ae370..583f5886f5a 100644 --- a/ui/litellm-dashboard/src/components/lens/traces/list/TracesTimeline.test.ts +++ b/ui/litellm-dashboard/src/components/lens/traces/list/TracesTimeline.test.ts @@ -1,6 +1,6 @@ import { describe, expect, it } from "vitest"; -import { bandForWindow, bucketRuns, dragUpdate, formatSpan } from "./TracesTimeline"; +import { bucketRuns } from "./TracesTimeline"; import type { TraceSummary } from "../types"; const HOUR = 3600 * 1000; @@ -41,43 +41,3 @@ describe("bucketRuns", () => { expect(buckets[5]).toMatchObject({ total: 1, failed: 0 }); }); }); - -describe("formatSpan", () => { - it("prints the largest two units, dropping zero parts", () => { - expect(formatSpan(45 * 60 * 1000)).toBe("45m"); - expect(formatSpan(6 * HOUR + 12 * 60 * 1000)).toBe("6h 12m"); - expect(formatSpan(24 * HOUR)).toBe("1d"); - expect(formatSpan(152 * 24 * HOUR + 23 * HOUR)).toBe("152d 23h"); - expect(formatSpan(-5)).toBe("0m"); - }); -}); - -describe("dragUpdate", () => { - const band = { lo: 10, hi: 14 }; - - it("selects between the press point and the pointer, in either direction", () => { - expect(dragUpdate({ mode: "select", origin: 20, band: { lo: 20, hi: 20 } }, 25)).toEqual({ lo: 20, hi: 25 }); - expect(dragUpdate({ mode: "select", origin: 20, band: { lo: 20, hi: 20 } }, 12)).toEqual({ lo: 12, hi: 20 }); - }); - - it("resizes one edge without letting it cross the other", () => { - expect(dragUpdate({ mode: "resize-lo", origin: 10, band }, 4)).toEqual({ lo: 4, hi: 14 }); - expect(dragUpdate({ mode: "resize-lo", origin: 10, band }, 30)).toEqual({ lo: 14, hi: 14 }); - expect(dragUpdate({ mode: "resize-hi", origin: 14, band }, 40)).toEqual({ lo: 10, hi: 40 }); - expect(dragUpdate({ mode: "resize-hi", origin: 14, band }, 2)).toEqual({ lo: 10, hi: 10 }); - }); - - it("pans the band keeping its width, clamped to the strip", () => { - expect(dragUpdate({ mode: "move", origin: 12, band }, 20)).toEqual({ lo: 18, hi: 22 }); - expect(dragUpdate({ mode: "move", origin: 12, band }, -50)).toEqual({ lo: 0, hi: 4 }); - expect(dragUpdate({ mode: "move", origin: 12, band }, 500, 60)).toEqual({ lo: 55, hi: 59 }); - }); -}); - -describe("bandForWindow", () => { - it("maps a selected window back to the buckets it covers", () => { - const buckets = bucketRuns([], range, 10); - expect(bandForWindow(buckets, { startMs: START + 2 * HOUR, endMs: START + 5 * HOUR })).toEqual({ lo: 2, hi: 4 }); - expect(bandForWindow(buckets, null)).toBeNull(); - }); -}); diff --git a/ui/litellm-dashboard/src/components/lens/traces/list/TracesTimeline.tsx b/ui/litellm-dashboard/src/components/lens/traces/list/TracesTimeline.tsx index 0cc580152fe..9cf947719d3 100644 --- a/ui/litellm-dashboard/src/components/lens/traces/list/TracesTimeline.tsx +++ b/ui/litellm-dashboard/src/components/lens/traces/list/TracesTimeline.tsx @@ -1,36 +1,15 @@ "use client"; -import { X } from "lucide-react"; import moment from "moment"; -import { useMemo, useRef, useState, type RefObject } from "react"; -import { useResizeObserver } from "usehooks-ts"; +import { useMemo } from "react"; -import { DotFieldCanvas, DotFieldRoot } from "@/components/shared/dotField/DotField"; -import type { DotBand, DotColumn } from "@/components/shared/dotField/dots"; -import { cn } from "@/lib/cva.config"; +import { Timeline, TIMELINE_BUCKETS, type Bucket, type TimeWindow } from "@/components/shared/timeline/Timeline"; import type { TraceSummary } from "../types"; import { traceAgentNames } from "../utils"; -const BUCKETS = 60; -const TICKS = 6; -const MINUTE_MS = 60 * 1000; -const HOUR_MS = 60 * MINUTE_MS; -const DAY_MS = 24 * HOUR_MS; -const EDGE_FORMAT = "MMM DD, HH:mm"; - -export interface TimeWindow { - startMs: number; - endMs: number; -} - -export interface Bucket extends DotColumn { - startMs: number; - endMs: number; -} - /** Run counts per equal-width time bucket across the window; runs outside it are dropped. */ -export function bucketRuns(runs: readonly TraceSummary[], range: TimeWindow, buckets = BUCKETS): Bucket[] { +export function bucketRuns(runs: readonly TraceSummary[], range: TimeWindow, buckets = TIMELINE_BUCKETS): Bucket[] { const width = (range.endMs - range.startMs) / buckets; const placed = runs.map((run) => ({ index: Math.floor((moment(run.start_time).valueOf() - range.startMs) / width), @@ -49,28 +28,6 @@ export function bucketRuns(runs: readonly TraceSummary[], range: TimeWindow, buc }); } -/** Compact window length, Logfire-style: "45m", "6h 12m", "7d", "152d 23h". */ -export function formatSpan(ms: number): string { - const totalMinutes = Math.max(0, Math.round(ms / MINUTE_MS)); - const days = Math.floor(totalMinutes / (24 * 60)); - const hours = Math.floor((totalMinutes % (24 * 60)) / 60); - const minutes = totalMinutes % 60; - if (days > 0) return hours > 0 ? `${days}d ${hours}h` : `${days}d`; - if (hours > 0) return minutes > 0 ? `${hours}h ${minutes}m` : `${hours}h`; - return `${minutes}m`; -} - -const tickFormat = (range: TimeWindow): string => (range.endMs - range.startMs > 2 * DAY_MS ? EDGE_FORMAT : "HH:mm"); - -const pct = (value: number): string => `${value * 100}%`; - -/** Keep the first / last tick label inside the strip; center the rest on their tick. */ -const tickShift = (t: number): string => { - if (t === 0) return "translateX(0)"; - if (t === 1) return "translateX(-100%)"; - return "translateX(-50%)"; -}; - interface TracesTimelineProps { runs: readonly TraceSummary[]; range: TimeWindow; @@ -78,273 +35,8 @@ interface TracesTimelineProps { onSelect: (selection: TimeWindow | null) => void; } -function BucketBar({ bucket }: { bucket: Bucket }) { - return
; -} - -function NowEdge() { - return ( - <> -